File size: 4,899 Bytes
d4eb935 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 | """Fixed-shape Qwen3 plus Kev pointer-head export wrapper."""
from __future__ import annotations
import math
import torch
from torch import nn
from torch.nn import functional as F
MASK_VALUE = -1e4
def rotate_half(value: torch.Tensor) -> torch.Tensor:
first, second = value.chunk(2, dim=-1)
return torch.cat((-second, first), dim=-1)
class KevExport(nn.Module):
"""Run one typed question against one state in a fixed token bucket.
Inputs use one-hot maps for the dynamic readout positions so Core ML never needs
data-dependent gather indices.
"""
def __init__(self, decision_model: nn.Module, length: int, max_options: int):
super().__init__()
language_model = decision_model.lm
config = language_model.config
self.embed_tokens = language_model.embed_tokens
self.layers = language_model.layers
self.norm = language_model.norm
self.pointer_q = decision_model.head.q
self.pointer_k = decision_model.head.k
self.length = length
self.max_options = max_options
self.num_heads = config.num_attention_heads
self.num_key_value_heads = config.num_key_value_heads
self.num_key_value_groups = self.num_heads // self.num_key_value_heads
self.head_dim = config.head_dim
self.hidden_size = config.hidden_size
self.attention_scale = self.head_dim**-0.5
self.pointer_scale = 1 / math.sqrt(self.pointer_q.out_features)
position_ids = torch.arange(length).unsqueeze(0)
probe = torch.zeros(1, length, config.hidden_size)
cos, sin = language_model.rotary_emb(probe, position_ids)
self.register_buffer("position_cos", cos.unsqueeze(1).detach().clone())
self.register_buffer("position_sin", sin.unsqueeze(1).detach().clone())
causal = torch.full((length, length), MASK_VALUE)
causal = torch.triu(causal, diagonal=1)
self.register_buffer("causal_mask", causal.view(1, 1, length, length))
def _rms_norm(self, hidden: torch.Tensor, norm: nn.Module) -> torch.Tensor:
doubled = torch.cat((hidden, -hidden), dim=-1)
normalized = F.layer_norm(doubled, (self.hidden_size * 2,), eps=norm.variance_epsilon)
return normalized[..., : self.hidden_size] * norm.weight
def _attention(self, layer: nn.Module, hidden: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:
attention = layer.self_attn
residual = hidden
normalized = self._rms_norm(hidden, layer.input_layernorm)
queries = (
attention.q_proj(normalized)
.view(1, self.length, self.num_heads, self.head_dim)
.transpose(1, 2)
)
keys = (
attention.k_proj(normalized)
.view(1, self.length, self.num_key_value_heads, self.head_dim)
.transpose(1, 2)
)
values = (
attention.v_proj(normalized)
.view(1, self.length, self.num_key_value_heads, self.head_dim)
.transpose(1, 2)
)
queries = attention.q_norm(queries)
keys = attention.k_norm(keys)
queries = queries * self.position_cos + rotate_half(queries) * self.position_sin
keys = keys * self.position_cos + rotate_half(keys) * self.position_sin
keys = keys.repeat_interleave(self.num_key_value_groups, dim=1)
values = values.repeat_interleave(self.num_key_value_groups, dim=1)
weights = torch.matmul(queries, keys.transpose(2, 3)) * self.attention_scale + mask
weights = torch.softmax(weights, dim=-1, dtype=torch.float32).to(queries.dtype)
attended = torch.matmul(weights, values).transpose(1, 2).reshape(1, self.length, -1)
hidden = residual + attention.o_proj(attended)
residual = hidden
normalized = self._rms_norm(hidden, layer.post_attention_layernorm)
return residual + layer.mlp(normalized)
def forward(
self,
input_ids: torch.Tensor,
attention_mask: torch.Tensor,
decide_map: torch.Tensor,
option_map: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor]:
padding = (1.0 - attention_mask.float()).view(1, 1, 1, self.length) * MASK_VALUE
mask = self.causal_mask + padding
hidden = self.embed_tokens(input_ids.long())
for layer in self.layers:
hidden = self._attention(layer, hidden, mask)
hidden = self._rms_norm(hidden, self.norm).float()
decide = torch.matmul(decide_map, hidden)
options = torch.matmul(option_map, hidden)
query = self.pointer_q(decide)
keys = self.pointer_k(options)
logits = torch.matmul(keys, query.transpose(1, 2)).squeeze(-1) * self.pointer_scale
supplied = option_map.sum(-1)
logits = logits * supplied + (1.0 - supplied) * MASK_VALUE
return logits, torch.softmax(logits, dim=-1)
|