test: fix gsa in SE multi-GPU test

This commit is contained in:
2026-06-01 04:26:03 +00:00
parent 419112dd3e
commit 9ba051cf49

View File

@@ -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}")