test: fix gsa in SE multi-GPU test
This commit is contained in:
@@ -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}")
|
||||
|
||||
Reference in New Issue
Block a user