Expand description
The same-step γ-correction on primitives — forward and analytic backward.
§Same-step γ-correction on primitives — forward and analytic backward
Primitive (F) port of super::super::diag, plus the analytic backward
the recompute node needs. The forward is used by
super::serial_recalculated’s K5; the backward by
super::combined_backward. Both carry the same SISO fast path: at
mimo_rank == 1 the m × m Gram matrix is a scalar, so every matmul over
the mimo_rank axis degenerates into a 1×K×1 / 1×1×N GEMM and is
replaced by a reduction plus broadcast multiplies.
Forward (per (b, n, l, h)):
qk_dot[m_out, m_in] = Σ_r C[m_out, r] · B[m_in, r]
y_diag[m_out, p] = γ · Σ_{m_in} qk_dot[m_out, m_in] · V[m_in, p]Structs§
- Diag
Grads - Gradients produced by
y_diag_correction_backward— they_diagterm’s contribution tod_v,d_c,d_band the whole ofd_gamma(no other term consumesγ).
Functions§
- y_
diag_ correction - The γ-weighted same-step correction
y_diag, on primitives. - y_
diag_ correction_ backward - Analytic backward of
y_diag_correction. - y_
diag_ 🔒correction_ backward_ mimo - General MIMO backward.
- y_
diag_ 🔒correction_ backward_ siso - SISO (
mimo_rank == 1) backward. - y_
diag_ 🔒correction_ mimo - General MIMO
y_diagon primitives. - y_
diag_ 🔒correction_ siso - SISO (
mimo_rank == 1)y_diagon primitives: astate_rankreduction plus a per-(b, n, l, h)scalar broadcast.