Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 4 additions & 2 deletions tests/unit/common.py
Original file line number Diff line number Diff line change
Expand Up @@ -575,8 +575,10 @@ def wrapper(*args: Any, **kwargs: Any):
def reduce_boolean_flags(flag: bool, op=all) -> bool:
if not dist.is_initialized():
return flag
device = get_accelerator().current_device()
tensor_flag = torch.tensor(1 if flag else 0, dtype=torch.int, device=device)
# current_device() is a rank id on CPU, not a valid torch device; use the device name.
device = get_accelerator().current_device_name()

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P1 Badge Add the required sign-off trailer

This is a non-merge commit, but its message has no Signed-off-by trailer, so it violates the repository's mandatory commit requirement and may be rejected by DCO/CI checks. Recreate the commit with --signoff using the configured Git identity.

AGENTS.md reference: AGENTS.md:L8-L8

Useful? React with 👍 / 👎.

# gloo rejects 0-dim inputs to all_gather_into_tensor, so carry the flag in a 1-dim tensor.
tensor_flag = torch.tensor([1 if flag else 0], dtype=torch.int, device=device)
world_size = dist.get_world_size()
tensor_flag_buf = torch.zeros(world_size, dtype=torch.int, device=device)
dist.all_gather_into_tensor(tensor_flag_buf, tensor_flag)
Expand Down
33 changes: 18 additions & 15 deletions tests/unit/v1/autotp/test_autotp_training.py
Original file line number Diff line number Diff line change
Expand Up @@ -424,7 +424,7 @@ def process_linear_layer(hidden_dim, input, output_dim=None):
torch_linear = nn.Linear(hidden_dim,
output_dim,
dtype=preferred_dtype(),
device=get_accelerator().current_device())
device=get_accelerator().current_device_name())
torch_out = torch_linear(input)
torch_loss = torch_out.sum()
torch_loss.backward()
Expand Down Expand Up @@ -488,7 +488,7 @@ def run_tp_layer_fwd_bwd(tp_size,
hidden_dim,
dtype=preferred_dtype(),
requires_grad=True,
device=get_accelerator().current_device())
device=get_accelerator().current_device_name())
dist.broadcast(input, groups.get_tensor_model_parallel_src_rank(), group=groups.get_tensor_model_parallel_group())

# Note: correctness checks below use standalone TP wrappers and do not
Expand All @@ -498,7 +498,7 @@ def run_tp_layer_fwd_bwd(tp_size,
linear = LinearLayer(deepcopy(torch_linear),
groups.get_tensor_model_parallel_group(),
gather_output=gather_output)
out = linear(input.to(get_accelerator().current_device()))
out = linear(input.to(get_accelerator().current_device_name()))
loss = out.sum()
loss.backward()

Expand All @@ -512,35 +512,35 @@ def run_tp_layer_fwd_bwd(tp_size,
torch_bias_grad = torch_linear.bias.grad.split(output_partition_sizes, dim=0)[tp_rank]

torch.testing.assert_close(linear.bias.grad,
torch_bias_grad.to(get_accelerator().current_device()),
torch_bias_grad.to(get_accelerator().current_device_name()),
atol=1e-3,
rtol=1e-3)
torch.testing.assert_close(linear.weight.grad,
torch_grad.to(get_accelerator().current_device()),
torch_grad.to(get_accelerator().current_device_name()),
atol=1e-3,
rtol=1e-3)
torch.testing.assert_close(expected_out.to(get_accelerator().current_device()).contiguous(),
torch.testing.assert_close(expected_out.to(get_accelerator().current_device_name()).contiguous(),
out.contiguous(),
atol=1e-2,
rtol=1e-2)
else:
linear = LinearAllreduce(deepcopy(torch_linear), groups.get_tensor_model_parallel_group())
input_ = torch.chunk(input, tp_size, dim=-1)[groups.get_tensor_model_parallel_rank()]
out = linear(input_.to(get_accelerator().current_device()))
out = linear(input_.to(get_accelerator().current_device_name()))
loss = out.sum()
loss.backward()

torch_grad = torch.chunk(torch_linear.weight.grad, tp_size, dim=1)[groups.get_tensor_model_parallel_rank()]
torch_bias_grad = torch_linear.bias.grad
torch.testing.assert_close(linear.bias.grad,
torch_bias_grad.to(get_accelerator().current_device()),
torch_bias_grad.to(get_accelerator().current_device_name()),
atol=1e-3,
rtol=1e-3)
torch.testing.assert_close(linear.weight.grad,
torch_grad.to(get_accelerator().current_device()),
torch_grad.to(get_accelerator().current_device_name()),
atol=1e-3,
rtol=1e-3)
torch.testing.assert_close(out, torch_out.to(get_accelerator().current_device()), atol=1e-2, rtol=1e-2)
torch.testing.assert_close(out, torch_out.to(get_accelerator().current_device_name()), atol=1e-2, rtol=1e-2)


@pytest.mark.sequential
Expand Down Expand Up @@ -577,7 +577,7 @@ def test_standard_cross_entropy_matches_unsharded_reference(self):
set_autotp_mode(training=True)
try:
torch.manual_seed(20260825)
model = UnevenVocabOutputModel(hidden_dim, vocab_size).to(get_accelerator().current_device())
model = UnevenVocabOutputModel(hidden_dim, vocab_size).to(get_accelerator().current_device_name())
reference_model = deepcopy(model)

autotp = AutoTP(module=model,
Expand All @@ -596,13 +596,16 @@ def test_standard_cross_entropy_matches_unsharded_reference(self):
assert model.lm_head.gather_output

torch.manual_seed(17)
reference_input = torch.randn(4, hidden_dim, device=get_accelerator().current_device(), requires_grad=True)
reference_input = torch.randn(4,
hidden_dim,
device=get_accelerator().current_device_name(),
requires_grad=True)
dist.broadcast(reference_input, src=0, group=tp_group)
tp_input = reference_input.detach().clone().requires_grad_(True)

partition_sizes = get_shard_size_list(vocab_size, self.world_size, model.lm_head.tp_meta, "lm_head")
labels = torch.tensor([0, vocab_size - 1, partition_sizes[0], 1],
device=get_accelerator().current_device())
device=get_accelerator().current_device_name())
reference_logits = reference_model(reference_input)
tp_logits = model(tp_input)
reference_loss = nn.functional.cross_entropy(reference_logits, labels)
Expand Down Expand Up @@ -693,7 +696,7 @@ def test(self, layer_type):
if is_model_parallel_parameter(param):
param.gather_params([param])

torch_linear = torch_linear.to(get_accelerator().current_device())
torch_linear = torch_linear.to(get_accelerator().current_device_name())
is_same_weights = all(
torch.equal(param1, param2) for param1, param2 in zip(tp_layer.parameters(), torch_linear.parameters()))

Expand Down Expand Up @@ -759,7 +762,7 @@ def test_uneven_linear_gather_params(self):
if is_model_parallel_parameter(param):
param.gather_params([param])

torch_linear = torch_linear.to(get_accelerator().current_device())
torch_linear = torch_linear.to(get_accelerator().current_device_name())
is_same_weights = all(
torch.equal(param1, param2) for param1, param2 in zip(tp_layer.parameters(), torch_linear.parameters()))

Expand Down
Loading