diff --git a/single_shot_inference.py b/single_shot_inference.py index afc749ef..5a1c1a80 100644 --- a/single_shot_inference.py +++ b/single_shot_inference.py @@ -72,92 +72,69 @@ def sinkhorn_knopp(M, t_max=20, eps=1e-6): class mHCBlock: - """One mHC block (attention or FFN). + """One mHC block — matches vLLM torch reference exactly. - Checkpoint weight mapping: - fn: (24, 28672) FP32 = stacked [W_pre(4,K); W_res(16,K); W_post(4,K)] - base: (24,) FP32 = bias, split as [S_pre(4); S_res(16); S_post(4)] - scale: (3,) FP32 = [alpha_pre, alpha_res, alpha_post] + Checkpoint: fn (24,28672), base (24,), scale (3,) + Split: pre_mix (n_hc=4), post_mix (n_hc=4), comb_mix (n_hc*n_hc=16) """ - def __init__(self, hidden_dim=7168, n_hc=4, t_max=20, device='cuda'): - self.d = hidden_dim + def __init__(self, hidden_dim=7168, n_hc=4, sinkhorn_repeat=20, device='cuda'): self.n_hc = n_hc - self.K = n_hc * hidden_dim # 28672 - self.t_max = t_max + self.K = n_hc * hidden_dim + self.sinkhorn_repeat = sinkhorn_repeat self.device = device - self.W_stacked = None # (24, K) FP32 - self.bias = None # (24,) FP32 - self.alphas = None # (3,) FP32 + self.fn = None + self.hc_scale = None + self.hc_base = None + self.rms_eps = 1e-6 + self.hc_pre_eps = 0.0 + self.hc_sinkhorn_eps = 1e-6 + self.hc_post_mult_value = 2.0 def load_from_checkpoint(self, fn, base, scale): - """Load from checkpoint tensors. All on target device, FP32.""" - self.W_stacked = fn.to(device=self.device, dtype=torch.float32).contiguous() - self.bias = base.to(device=self.device, dtype=torch.float32).contiguous() - self.alphas = scale.to(device=self.device, dtype=torch.float32).contiguous() + self.fn = fn.to(device=self.device, dtype=torch.float32).contiguous() + self.hc_base = base.to(device=self.device, dtype=torch.float32).contiguous() + self.hc_scale = scale.to(device=self.device, dtype=torch.float32).contiguous() - def _dynamic_params(self, X_l): - """Compute A_l, B_l, C_l from residual state. + def pre_block(self, residual): + """residual: (T, n_hc, d) BF16 → layer_input (T, d), ctx""" + n = self.n_hc + K = self.K + eps = self.rms_eps + T = residual.shape[0] - X_l: (T, n_hc, d) BF16 - Returns: A_l (T, n_hc), B_l (T, n_hc, n_hc) FP32, C_l (T, n_hc) - """ - T, n, d = X_l.shape - n_hc = self.n_hc + res_flat = residual.reshape(T, K).float() # (T, K) - # Flatten and project with RMSNorm - X_flat = X_l.reshape(T, self.K) # (T, K) BF16 + # Project with RMSNorm on input + mixes = torch.matmul(res_flat, self.fn.t()) # (T, 24) + sqrsum = res_flat.square().sum(dim=-1, keepdim=True) + mixes = mixes * torch.rsqrt(sqrsum / K + eps) - # RMSNorm - x_f32 = X_flat.float() - rms = x_f32.pow(2).mean(-1, keepdim=True).add(1e-6).rsqrt() - x_normed = x_f32 * rms # (T, K) FP32 + # Split + pre_logits = mixes[:, :n] * self.hc_scale[0] + self.hc_base[:n] + pre_mix = torch.sigmoid(pre_logits) + self.hc_pre_eps # (T, 4) - # Project: (T, K) @ (24, K)^T → (T, 24) - proj = torch.nn.functional.linear(x_normed.bfloat16(), self.W_stacked.bfloat16()).float() - proj = proj + self.bias.unsqueeze(0) # add bias + post_logits = mixes[:, n:2*n] * self.hc_scale[1] + self.hc_base[n:2*n] + post_mix = torch.sigmoid(post_logits) * self.hc_post_mult_value # (T, 4) - # Split into A, B, C - i0, i1, i2 = n_hc, n_hc + n_hc * n_hc, 24 - A_raw = proj[:, :i0] # (T, 4) - B_raw = proj[:, i0:i1] # (T, 16) - C_raw = proj[:, i1:i2] # (T, 4) + comb_logits = (mixes[:, 2*n:].reshape(T, n, n) * self.hc_scale[2] + + self.hc_base[2*n:].reshape(1, n, n)) + comb_mix = torch.softmax(comb_logits, dim=-1) + self.hc_sinkhorn_eps + comb_mix = comb_mix / (comb_mix.sum(dim=-2, keepdim=True) + self.hc_sinkhorn_eps) + for _ in range(self.sinkhorn_repeat - 1): + comb_mix = comb_mix / (comb_mix.sum(dim=-1, keepdim=True) + self.hc_sinkhorn_eps) + comb_mix = comb_mix / (comb_mix.sum(dim=-2, keepdim=True) + self.hc_sinkhorn_eps) - # Split bias into S_pre, S_res, S_post - S_pre = self.bias[:n_hc] - S_res = self.bias[n_hc:n_hc + n_hc * n_hc] - S_post = self.bias[n_hc + n_hc * n_hc:] + # layer_input = sum(pre_mix * residual) + layer_input = (pre_mix.unsqueeze(-1) * residual.float()).sum(dim=1).bfloat16() - # Apply gating + biases - a_pre, a_res, a_post = self.alphas[0], self.alphas[1], self.alphas[2] - A_tilde = a_pre * A_raw + S_pre.unsqueeze(0) - B_tilde = a_res * B_raw + S_res.unsqueeze(0) - C_tilde = a_post * C_raw + S_post.unsqueeze(0) - - # Constraints - A_l = torch.sigmoid(A_tilde).bfloat16() # (T, 4) ∈ (0,1) - C_l = (2.0 * torch.sigmoid(C_tilde)).bfloat16() # (T, 4) ∈ (0,2) - B_exp = torch.exp(B_tilde).reshape(T, n_hc, n_hc) # (T, 4, 4) - B_l = sinkhorn_knopp(B_exp, self.t_max) # FP32, doubly stochastic - - return A_l, B_l, C_l + return layer_input, (post_mix, comb_mix) - def pre_block(self, X_l): - """X_l: (T, n_hc, d) → x_in: (T, d), ctx""" - A_l, B_l, C_l = self._dynamic_params(X_l) - # x_in = A_l @ X_l: (T, 1, n_hc) bmm (T, n_hc, d) → (T, 1, d) → (T, d) - x_in = torch.bmm(A_l.unsqueeze(1).float(), X_l.float()).squeeze(1).bfloat16() - return x_in, (B_l, C_l) - - def post_block(self, X_l, F_out, ctx): - """X_l: (T, n_hc, d), F_out: (T, d) → X_next: (T, n_hc, d)""" - B_l, C_l = ctx - # X_next = B_l @ X_l + C_l ⊗ F_out - BX = torch.bmm(B_l, X_l.float()) # (T, n_hc, d) FP32 - CF = C_l.unsqueeze(-1).float() * F_out.unsqueeze(1).float() # (T, n_hc, d) FP32 - return (BX + CF).bfloat16() - - -# ===================================================================== + def post_block(self, residual, F_out, ctx): + """residual: (T, n_hc, d), F_out: (T, d) → (T, n_hc, d)""" + post_mix, comb_mix = ctx + mixed_residual = torch.einsum('tij,tjh->tjh', comb_mix, residual.float()) + post_term = post_mix.unsqueeze(-1) * F_out.unsqueeze(1).float() + return (mixed_residual + post_term).bfloat16() # RoPE # =====================================================================