diff --git a/test_se_multi_gpu.py b/test_se_multi_gpu.py index 794c8f55..edb40c43 100644 --- a/test_se_multi_gpu.py +++ b/test_se_multi_gpu.py @@ -56,13 +56,15 @@ for gpu in [0, 1]: # Run x = torch.randn(1, 7168, dtype=torch.bfloat16, device=dev) - out = se.run(x) - # Fix gsa - if hasattr(se, '_saved_l1_gsa'): - se._l1_activation_global_scale = se._saved_l1_gsa - if hasattr(se, '_saved_l2_gsa'): - se._l2_activation_global_scale = se._saved_l2_gsa + # Must set gsa AFTER _ensure_initialized but BEFORE run + # _ensure_initialized is called lazily in run(), so we need to call it first + se._ensure_initialized() + # Now fix the gsa + se._l1_activation_global_scale = gisc.float().item() + se._l2_activation_global_scale = disc.float().item() + + out = se.run(x) has_nan = torch.isnan(out).any().item() print(f"GPU {gpu}: |out|={out.abs().max().item() if not has_nan else 'NaN'} has_nan={has_nan} shape={out.shape}")