fix quality formatting

This commit is contained in:
Pepijn
2026-07-15 18:27:13 +02:00
parent ccecdbc769
commit 55d9ff740e
2 changed files with 13 additions and 21 deletions
+1 -3
View File
@@ -31,9 +31,7 @@ def test_message_recipe_validates_unknown_binding():
def test_canonical_recipe_loads(): def test_canonical_recipe_loads():
"""The canonical PI052 blend YAML loads + validates.""" """The canonical PI052 blend YAML loads + validates."""
recipe = TrainingRecipe.from_yaml( recipe = TrainingRecipe.from_yaml(Path("src/lerobot/configs/recipes/subtask_mem_vqa_speech.yaml"))
Path("src/lerobot/configs/recipes/subtask_mem_vqa_speech.yaml")
)
assert recipe.blend is not None assert recipe.blend is not None
assert sum(c.weight for c in recipe.blend.values()) == pytest.approx(1.0) assert sum(c.weight for c in recipe.blend.values()) == pytest.approx(1.0)
@@ -105,12 +105,8 @@ def test_sdpa_parity_with_eager_block_bidirectional(num_heads, num_kv_heads, hea
module = _mock_self_attn(num_heads // num_kv_heads) module = _mock_self_attn(num_heads // num_kv_heads)
out_eager, _ = modeling_gemma.eager_attention_forward( out_eager, _ = modeling_gemma.eager_attention_forward(module, q, k, v, mask, scaling)
module, q, k, v, mask, scaling out_sdpa, _ = sdpa_attention_forward(module, q, k, v, mask, scaling)
)
out_sdpa, _ = sdpa_attention_forward(
module, q, k, v, mask, scaling
)
assert out_eager.shape == out_sdpa.shape assert out_eager.shape == out_sdpa.shape
torch.testing.assert_close(out_sdpa, out_eager, atol=1e-5, rtol=1e-4) torch.testing.assert_close(out_sdpa, out_eager, atol=1e-5, rtol=1e-4)
@@ -123,12 +119,8 @@ def test_sdpa_parity_bf16():
mask = _block_bidirectional_mask(bsize, seq_len, [5, 6, 6], torch.bfloat16) mask = _block_bidirectional_mask(bsize, seq_len, [5, 6, 6], torch.bfloat16)
module = _mock_self_attn(num_heads // num_kv_heads) module = _mock_self_attn(num_heads // num_kv_heads)
out_eager, _ = modeling_gemma.eager_attention_forward( out_eager, _ = modeling_gemma.eager_attention_forward(module, q, k, v, mask, scaling)
module, q, k, v, mask, scaling out_sdpa, _ = sdpa_attention_forward(module, q, k, v, mask, scaling)
)
out_sdpa, _ = sdpa_attention_forward(
module, q, k, v, mask, scaling
)
torch.testing.assert_close(out_sdpa, out_eager, atol=2e-2, rtol=2e-2) torch.testing.assert_close(out_sdpa, out_eager, atol=2e-2, rtol=2e-2)
@@ -138,7 +130,9 @@ def test_sdpa_parity_backward():
bsize, num_heads, num_kv_heads, seq_len, head_dim = 1, 4, 2, 9, 32 bsize, num_heads, num_kv_heads, seq_len, head_dim = 1, 4, 2, 9, 32
scaling = head_dim**-0.5 scaling = head_dim**-0.5
q, k, v = _build_inputs(bsize, num_heads, num_kv_heads, seq_len, head_dim, torch.float32) q, k, v = _build_inputs(bsize, num_heads, num_kv_heads, seq_len, head_dim, torch.float32)
q.requires_grad_(True); k.requires_grad_(True); v.requires_grad_(True) q.requires_grad_(True)
k.requires_grad_(True)
v.requires_grad_(True)
mask = _block_bidirectional_mask(bsize, seq_len, [3, 3, 3], torch.float32) mask = _block_bidirectional_mask(bsize, seq_len, [3, 3, 3], torch.float32)
module = _mock_self_attn(num_heads // num_kv_heads) module = _mock_self_attn(num_heads // num_kv_heads)