""" Ternary Quantized Diagonal State-Space Model (Parallel) """ import torch import torch.nn as nn import torch.nn.functional as F from .bitlinear import BitLinear, RMSNorm class Q88Quantize(torch.autograd.Function): """Q8.8 fixed-point quantization with straight-through estimator.""" @staticmethod def forward(ctx, x): """ Quantize to Q8.8 format (8 integer bits, 8 fractional bits) Range: [-128, 127.99609375] """ scale = 2**8 # 256 # Quantize: scale -> round -> clamp to int16 range -> dequantize x_scaled = x * scale x_int = torch.clamp(torch.round(x_scaled), -32768, 32767) x_quant = x_int / scale return x_quant @staticmethod def backward(ctx, grad_output): # Straight-through estimator: pass gradients unchanged return grad_output class SSMBlock(nn.Module): """ Diagonal (convolutional) SSM Block with ternary BitLinear projections. Architecture: Input → B projection → diagonal SSM convolution → C projection → Output State dynamics (training, parallel): s_t = sum_{k=0}^t (B x_k) State dynamics (inference, step-wise): s_t = s_{t-1} + B x_t y_t = C s_t """ def __init__(self, config): super().__init__() self.config = config self.d_model = config.d_model self.d_state = config.d_state # ===================================================================== # Stationary ternary projections # ===================================================================== self.b_proj = BitLinear(self.d_model, self.d_state, bias=False) self.c_proj = BitLinear(self.d_state, self.d_model, bias=False) # A matrix: identity matrix scaled by a single scalar decay factor self.register_buffer("a_log", torch.log(torch.tensor(0.9))) self.dropout = nn.Dropout(config.dropout) # --------------------------------------------------------------------- # Training / parallel forward # --------------------------------------------------------------------- def forward(self, x, mask=None): """ Args: x: [batch, seq_len, d_model] mask: unused (SSM is causal by construction) Returns: y: [batch, seq_len, d_model] """ B, L, _ = x.shape # Input projection u = self.b_proj(x) # Compute decay with Q8.8 quantization decay = torch.exp(self.a_log) # scalar decay_quant = decay + (Q88Quantize.apply(decay) - decay).detach() L = u.size(1) device = u.device dtype = u.dtype # Decay powers with quantization: [L, 1] (broadcasts across d_state) t = torch.arange(L, device=device, dtype=dtype).unsqueeze(1) # [L, 1] decay_pows = decay_quant ** t # [L, 1] #decay_pows = decay_pows + (Q88Quantize.apply(decay_pows) - decay_pows).detach() inv_decay_pows = decay_pows.reciprocal() #inv_decay_pows = inv_decay_pows + (Q88Quantize.apply(inv_decay_pows) - inv_decay_pows).detach() # Reweight, cumsum, reweight back s = torch.cumsum(u * inv_decay_pows.unsqueeze(0), dim=1) # [B, L, d_state] s = s * decay_pows.unsqueeze(0) # Output projection y = self.c_proj(s) y = self.dropout(y) return y # --------------------------------------------------------------------- # Autoregressive single-step inference # --------------------------------------------------------------------- def step(self, x, state): """ Single timestep SSM update (for autoregressive decoding). Args: x: [batch, d_model] state: [batch, d_state] Returns: output: [batch, d_model] new_state: [batch, d_state] """ decay = torch.exp(self.a_log) # scalar new_state = decay * state + self.b_proj(x) # [batch, d_state] output = self.c_proj(new_state) return output, new_state # --------------------------------------------------------------------- # State utilities # --------------------------------------------------------------------- def init_state(self, batch_size, device, dtype): """Initialize hidden state.""" return torch.zeros(batch_size, self.d_state, device=device, dtype=dtype) # --------------------------------------------------------------------- # Export parameters for inference / FPGA # --------------------------------------------------------------------- def get_inference_params(self): """ Export parameters for deployment. Returns: dict with quantized projections and diagonal A """ with torch.no_grad(): return { "b_proj": self.b_proj.get_inference_params(), "c_proj": self.c_proj.get_inference_params(), }