diff --git a/tests/pytorch/test_fused_optimizer.py b/tests/pytorch/test_fused_optimizer.py index a2863cba98..6832ef89dd 100644 --- a/tests/pytorch/test_fused_optimizer.py +++ b/tests/pytorch/test_fused_optimizer.py @@ -166,6 +166,19 @@ def test_frozen_model(self): torch.testing.assert_close(ref_param, tst_param) + def test_empty_param_at_end_of_group(self): + tensors = [ + torch.ones(4, dtype=torch.float, device="cuda"), + torch.empty(0, dtype=torch.float, device="cuda"), + ] + ref_param, tst_param, ref_optim, tst_optim = self.gen_param_optim(tensors, self.options) + + self.gen_grad(ref_param, tst_param) + ref_optim.step() + tst_optim.step() + + torch.testing.assert_close(ref_param, tst_param) + def gen_precision_aware_test( self, use_fp8_params, @@ -796,6 +809,19 @@ def test_float(self): def test_half(self): self.gen_single_type_test(param_type=torch.float16) + def test_empty_param_at_end_of_group(self): + tensors = [ + torch.ones(4, dtype=torch.float, device="cuda"), + torch.empty(0, dtype=torch.float, device="cuda"), + ] + ref_param, tst_param, ref_optim, tst_optim = self.gen_param_optim(tensors, self.options) + + self.gen_grad(ref_param, tst_param) + ref_optim.step() + tst_optim.step() + + torch.testing.assert_close(ref_param, tst_param) + class Model(torch.nn.Module): def __init__(self): diff --git a/transformer_engine/common/multi_tensor/multi_tensor_apply.cuh b/transformer_engine/common/multi_tensor/multi_tensor_apply.cuh index 6710a161b3..d6ac23d4af 100644 --- a/transformer_engine/common/multi_tensor/multi_tensor_apply.cuh +++ b/transformer_engine/common/multi_tensor/multi_tensor_apply.cuh @@ -75,6 +75,9 @@ void multi_tensor_apply(int64_t block_size, int64_t chunk_size, loc_tensor_info++; auto chunks_this_tensor = (tensor_lists[0][t]->numel() + chunk_size - 1) / chunk_size; + NVTE_CHECK(chunks_this_tensor > 0, + "multi_tensor_apply expects tensors with at least one chunk; zero-sized tensors " + "must be filtered before launch because they skip the chunk loop"); for (auto chunk = 0; chunk < chunks_this_tensor; chunk++) { tl.block_to_tensor[loc_block_info] = loc_tensor_info - 1; diff --git a/transformer_engine/pytorch/optimizers/fused_sgd.py b/transformer_engine/pytorch/optimizers/fused_sgd.py index d7ab3fe9fe..10151e1406 100644 --- a/transformer_engine/pytorch/optimizers/fused_sgd.py +++ b/transformer_engine/pytorch/optimizers/fused_sgd.py @@ -295,20 +295,19 @@ def step(self, closure=None): for _, (launch_set, first_run) in enumerate(zip(launch_sets, first_runs)): assert len(launch_set[0]) == len(launch_set[1]) assert len(launch_set[0]) == len(launch_set[2]) - if len(launch_set[0]) > 0: - multi_tensor_applier( - self.multi_tensor_sgd, - self._dummy_overflow_buf, - launch_set, - weight_decay, - momentum, - dampening, - group["lr"], - nesterov, - first_run, - self.wd_after_momentum, - 1.0 / self.most_recent_scale, - ) + multi_tensor_applier( + self.multi_tensor_sgd, + self._dummy_overflow_buf, + launch_set, + weight_decay, + momentum, + dampening, + group["lr"], + nesterov, + first_run, + self.wd_after_momentum, + 1.0 / self.most_recent_scale, + ) self.most_recent_scale = 1.0 self.scale_set_by_backward = False diff --git a/transformer_engine/pytorch/optimizers/multi_tensor_apply.py b/transformer_engine/pytorch/optimizers/multi_tensor_apply.py index a5cbd27337..e7d4fa0db3 100644 --- a/transformer_engine/pytorch/optimizers/multi_tensor_apply.py +++ b/transformer_engine/pytorch/optimizers/multi_tensor_apply.py @@ -18,6 +18,16 @@ def __call__(self, op, noop_flag_buffer, tensor_lists, *args): if isinstance(t, DTensor): tensor_lists[i][j] = t._local_tensor + if any(len(tensors) != len(tensor_lists[0]) for tensors in tensor_lists): + raise RuntimeError("Expected aligned multi-tensor lists.") + + keep_slot = [tensor.numel() > 0 for tensor in tensor_lists[0]] + for i, tensors in enumerate(tensor_lists): + tensor_lists[i] = [tensor for tensor, keep in zip(tensors, keep_slot) if keep] + + if not tensor_lists[0]: + return None + return op(self.chunk_size, noop_flag_buffer, tensor_lists, *args)