"""Resident Bend encoder projections and optional decision-head normalization. The packed FP32 kernels own matrix arithmetic and coarse parallel scheduling. PyTorch supplies embedding lookup, SDPA, activations and the selected-output head. """ import numpy as np import torch from torch import nn class BendLayerNorm(nn.Module): def __init__(self, original, reducer): super().__init__() self.weight, self.bias = original.weight, original.bias self.eps = original.eps self.normalized_shape = original.normalized_shape self.reducer = reducer def forward(self, x): if x.device.type != 'cpu' or x.dtype != torch.float32: raise ValueError('Bend LayerNorm requires CPU float32; use torch backend for CUDA') if torch.is_grad_enabled(): raise RuntimeError('Bend transformer kernels require torch.inference_mode()') width = x.shape[-1] values = x.detach().contiguous().numpy().reshape(-1, width) gamma = self.weight.detach().numpy() if self.weight is not None else None beta = self.bias.detach().numpy() if self.bias is not None else None parts = [self.reducer.layernorm(values[start:start + 1048576 // width], gamma, beta, self.eps) for start in range(0, len(values), 1048576 // width)] result = parts[0] if len(parts) == 1 else np.concatenate(parts) return torch.from_numpy(result.reshape(x.shape)) class BendHeadLayer(nn.Module): """Explicit pre-norm block; prevents fused MHA from bypassing Bend norms.""" def __init__(self, layer, reducer): super().__init__() if not layer.norm_first: raise ValueError('Julia expects pre-norm transformer layers') self.layer = layer self.norm1 = BendLayerNorm(layer.norm1, reducer) self.norm2 = BendLayerNorm(layer.norm2, reducer) def forward(self, x, src_key_padding_mask=None): normalized = self.norm1(x) attention = self.layer.self_attn(normalized, normalized, normalized, key_padding_mask=src_key_padding_mask, need_weights=False)[0] x = x + self.layer.dropout1(attention) hidden = self.layer.linear2(self.layer.dropout( self.layer.activation(self.layer.linear1(self.norm2(x))))) return x + self.layer.dropout2(hidden) def install_bend_head(model, reducer): if model.head is not None: model.head.layers = nn.ModuleList([BendHeadLayer(layer, reducer) for layer in model.head.layers]) model.scorer[0] = BendLayerNorm(model.scorer[0], reducer) model.eval() class BendLinear(nn.Module): """Inference-only CPU projection using resident weights and Bend arithmetic.""" def __init__(self, original, reducer): super().__init__() from .native import BendMatrix self.weight, self.bias = original.weight, original.bias self.in_features, self.out_features = original.in_features, original.out_features self.matrix = BendMatrix(original.weight.detach().numpy(), reducer) self._weight_version = self.weight._version self._weight_pointer = self.weight.data_ptr() def _check_weight(self): if self.weight._version != self._weight_version or self.weight.data_ptr() != self._weight_pointer: raise RuntimeError('Bend resident weights changed; recreate the inference backend') def forward(self, x): if x.device.type != 'cpu' or x.dtype != torch.float32 or torch.is_grad_enabled(): raise ValueError('Bend dense projections require CPU FP32 inference_mode') self._check_weight() result = self.matrix.tensor(x, validate=False) return result if self.bias is None else result + self.bias def install_bend_encoder(model, reducer): """Route all encoder linear projections through Bend; attention remains SDPA.""" count = 0 def install(module): nonlocal count for name, child in list(module.named_children()): if type(child) is nn.Linear: setattr(module, name, BendLinear(child, reducer)) count += 1 else: install(child) install(model.encoder) model.eval() return count