Download source/export_model.py from FluidInference/kev-0.6b-coreml: direct link, hf CLI and curl.
- Browser
- Download file 4.9 kB
-
https://huggingface.co/FluidInference/kev-0.6b-coreml/resolve/main/source/export_model.py
- Command line
-
hf download hf://FluidInference/kev-0.6b-coreml/source/export_model.py
-
curl -L -o export_model.py https://huggingface.co/FluidInference/kev-0.6b-coreml/resolve/main/source/export_model.py
4.9 kB
| """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) | |