diff --git a/ScaFFold/utils/losses.py b/ScaFFold/utils/losses.py index 0a1636d..07a386d 100644 --- a/ScaFFold/utils/losses.py +++ b/ScaFFold/utils/losses.py @@ -128,7 +128,7 @@ def compute_sharded_cross_entropy_loss( Each rank only sees a local spatial shard, so we cannot use the local `reduction="mean"` result directly. Instead we: - 1. compute the local CE numerator with `reduction="sum"`, + 1. compute the local CE numerator by summing the per-voxel losses, 2. build the correct global denominator, 3. all-reduce numerator and denominator together across the spatial mesh in a single collective, and @@ -146,40 +146,52 @@ def compute_sharded_cross_entropy_loss( autocast_device = device_type if device_type != "mps" else "cpu" with torch.autocast(autocast_device, enabled=False): - # Accumulate CE in full precision. Using reduction="sum" gives us the - # numerator of the final global mean; if class weights are present, - # PyTorch applies the target-class weight to each voxel here. When the - # caller already computed log-softmax, NLL over it is identical to - # cross-entropy over the raw logits but avoids a second full upcast. + # Accumulate CE in full precision. Summing the per-voxel losses gives + # us the numerator of the final global mean; if class weights are + # present, PyTorch applies the target-class weight to each voxel here. + # When the caller already computed log-softmax, NLL over it is + # identical to cross-entropy over the raw logits but avoids a second + # full upcast. + # + # reduction="none" plus .sum(), not reduction="sum": the fused CUDA + # reduction accumulates with atomicAdd, so its result depends on block + # retire order, and it has no deterministic implementation + # (`more_determinism` uses warn_only=True, so it would only warn). The + # separate .sum() is a fixed-order reduction. if log_probs is not None: - local_ce_sum = F.nll_loss( + local_ce = F.nll_loss( log_probs, local_labels, weight=class_weights, - reduction="sum", + reduction="none", ) else: - local_ce_sum = F.cross_entropy( + local_ce = F.cross_entropy( local_preds.float(), local_labels, weight=class_weights, - reduction="sum", + reduction="none", ) + local_ce_sum = local_ce.sum() + # Neither branch may read device memory from the host: a sync here + # drains the launch queue mid-step. set_sync_debug_mode("error") + # catches a regression. if class_weights is None: # Sum the actual local voxel counts across spatial shards. We use # an all-reduced count instead of numel()*num_shards because shard # sizes can differ at chunk boundaries. - local_normalizer = local_ce_sum.new_tensor(float(local_labels.numel())) + # + # new_full rather than new_tensor: the value travels as a kernel + # argument, not a pageable host-to-device copy. + local_normalizer = local_ce_sum.new_full((), float(local_labels.numel())) else: # Weighted CE divides by sum(weight[target_i]) over all voxels. - # Build that denominator from the local label histogram. - local_class_counts = torch.bincount( - local_labels.reshape(-1), minlength=class_weights.numel() - ).to(dtype=local_ce_sum.dtype) - local_normalizer = torch.dot( - local_class_counts, class_weights.to(dtype=local_ce_sum.dtype) - ) + # Not torch.bincount: it reads the largest label back to the host + # to size its output, even with minlength given. + local_normalizer = class_weights.to(dtype=local_ce_sum.dtype)[ + local_labels + ].sum() # Reduce the CE numerator and its denominator across the spatial shards in # one collective (they share the same mesh) rather than two, halving the diff --git a/tests/test_perf_hotpath.py b/tests/test_perf_hotpath.py index eac6a5f..a9a1c45 100644 --- a/tests/test_perf_hotpath.py +++ b/tests/test_perf_hotpath.py @@ -20,6 +20,7 @@ numerically equivalent to the straightforward reference it replaced. """ +import pytest import torch import torch.nn.functional as F @@ -67,6 +68,94 @@ def test_ce_log_probs_path_matches_cross_entropy(): assert torch.allclose(plain, ref, atol=1e-6) +@pytest.mark.gpu +def test_gpu_ce_numerator_is_bitwise_reproducible(): + # The CE numerator is summed outside the loss kernel: reduction="sum" on + # CUDA accumulates with atomicAdd and varies between identical calls. Both + # entry points (log_probs via NLL, raw logits via CE) are checked bitwise. + # + # The volume is load-bearing: the atomics only collide with many blocks in + # flight, so a smaller shape passes either way. 128**3 with 7 classes is + # scale 7 at the shipped n_categories. + torch.manual_seed(3) + device = torch.device("cuda") + b, c, n = 1, 7, 128 + preds = torch.randn(b, c, n, n, n, device=device) + labels = torch.randint(0, c, (b, n, n, n), device=device) + weights = torch.rand(c, device=device) + 0.5 + log_probs = F.log_softmax(preds.float(), dim=1) + + for w in (None, weights): + for kwargs in ({"log_probs": log_probs}, {}): + values = { + compute_sharded_cross_entropy_loss( + preds, labels, None, (1,), "cuda", w, **kwargs + ).item() + for _ in range(50) + } + assert len(values) == 1, f"{len(values)} distinct CE values: {values}" + + +@pytest.mark.gpu +def test_gpu_ce_does_not_synchronize(): + # The CE term issues no host-device sync; one would drain the launch queue + # mid-step and is invisible in the loss value. set_sync_debug_mode does + # not catch every synchronizing op, so this is a floor, not a proof. The + # shape only has to reach both normalizer branches and both entry points. + torch.manual_seed(5) + device = torch.device("cuda") + b, c, n = 1, 7, 64 + preds = torch.randn(b, c, n, n, n, device=device) + labels = torch.randint(0, c, (b, n, n, n), device=device) + weights = torch.rand(c, device=device) + 0.5 + log_probs = F.log_softmax(preds.float(), dim=1) + + # .item() is itself a sync: inspect the results after the mode is off. + outs = [] + torch.cuda.synchronize() + torch.cuda.set_sync_debug_mode("error") + try: + for w in (weights, None): + for kwargs in ({"log_probs": log_probs}, {}): + outs.append( + compute_sharded_cross_entropy_loss( + preds, labels, None, (1,), "cuda", w, **kwargs + ) + ) + finally: + torch.cuda.set_sync_debug_mode("default") + + torch.cuda.synchronize() + assert all(torch.isfinite(o) for o in outs) + + +@pytest.mark.gpu +def test_gpu_ce_survives_strict_deterministic_algorithms(): + # The CE path is legal under strict determinism. more_determinism uses + # warn_only=True, so a nondeterministic kernel would only warn there; + # strict mode raises, since the fused loss reduction has no deterministic + # implementation. + torch.manual_seed(4) + device = torch.device("cuda") + b, c = 1, 5 + preds = torch.randn(b, c, 16, 16, 16, device=device) + labels = torch.randint(0, c, (b, 16, 16, 16), device=device) + weights = torch.rand(c, device=device) + 0.5 + log_probs = F.log_softmax(preds.float(), dim=1) + + was_deterministic = torch.are_deterministic_algorithms_enabled() + was_warn_only = torch.is_deterministic_algorithms_warn_only_enabled() + torch.use_deterministic_algorithms(True) + try: + for w in (None, weights): + for kwargs in ({"log_probs": log_probs}, {}): + compute_sharded_cross_entropy_loss( + preds, labels, None, (1,), "cuda", w, **kwargs + ) + finally: + torch.use_deterministic_algorithms(was_deterministic, warn_only=was_warn_only) + + def test_ce_uses_single_spatial_collective(monkeypatch): # The CE numerator and its normalizer are reduced together in one # SpatialAllReduce, not two. Count applications; the packed path issues