test: compare logits_positions rows with assert_close
- sliced M=2 and full M=5 lm_head projections may pick different GEMM kernels whose accumulation order differs in the last float32 bits, so torch.equal flakes by machine and thread count
This commit is contained in:
@@ -157,7 +157,9 @@ def test_forward_logits_positions_projects_only_requested_rows():
|
|||||||
|
|
||||||
assert full["logits"].shape == (5, config.vocab_size)
|
assert full["logits"].shape == (5, config.vocab_size)
|
||||||
assert sliced["logits"].shape == (2, config.vocab_size)
|
assert sliced["logits"].shape == (2, config.vocab_size)
|
||||||
assert torch.equal(sliced["logits"], full["logits"][last_rows])
|
# Projecting M=2 rows vs M=5 rows may pick different GEMM kernels and
|
||||||
|
# differ in the last float32 bits — assert_close, not bit equality.
|
||||||
|
torch.testing.assert_close(sliced["logits"], full["logits"][last_rows])
|
||||||
assert torch.equal(sliced["hidden_states"], full["hidden_states"][last_rows])
|
assert torch.equal(sliced["hidden_states"], full["hidden_states"][last_rows])
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user