Julia-1 / julia /router /transformer.py
kleeedolinux
Publish Julia 1 model and Python runtime
5278d6b
Raw
History Blame Contribute Delete
4.25 kB
"""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