From baee36e728111de98eb5126f1d05e9851f619404 Mon Sep 17 00:00:00 2001 From: biondizzle Date: Sun, 31 May 2026 20:23:18 +0000 Subject: [PATCH] Fix dtype mismatch in validate_layer: cast flat to float before F.linear --- tests/validate_layer.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/validate_layer.py b/tests/validate_layer.py index 2b5ea33e..a2071c45 100644 --- a/tests/validate_layer.py +++ b/tests/validate_layer.py @@ -210,7 +210,7 @@ def validate_layer(li, all_weights, cfg, device='cuda:0'): flat = (X_flat * rms_inv).to(torch.bfloat16) # F.linear + split [pre(4), post(4), comb(16)] - proj = torch.nn.functional.linear(flat, fn).float() + proj = torch.nn.functional.linear(flat.float(), fn).float() pre_w, post_w, comb_w = proj.split([n_hc, n_hc, n_hc * n_hc], dim=-1) pre_b, post_b, comb_b = base.split([n_hc, n_hc, n_hc * n_hc]) pre_s, post_s, comb_s = scale.unbind(0)