Qyvos / julia /model.py
Manusagents's picture
Qyvos v1: Julia-1 backbone (bit-exact) + Open-Jev head fine-tune (30k rows, low-RAM protocol)
31f7037 verified
Raw History Blame Contribute Delete
7.54 kB
"""Marker based decision model and checkpoint serialization."""
import json
from pathlib import Path
import torch
from torch import nn
from torch.nn import functional as F
from torch.utils.checkpoint import checkpoint
from safetensors.torch import load_file, save_file
from transformers import AutoConfig, AutoModel
class JuliaDecisionModel(nn.Module):
def __init__(self, encoder, head_layers=2, n_act=2, dropout=.1):
super().__init__()
self.encoder = encoder
width = encoder.config.hidden_size
self.settings = dict(head_layers=head_layers, n_act=n_act, dropout=dropout)
layer = nn.TransformerEncoderLayer(width, max(1, width // 64), 4 * width,
dropout, batch_first=True, norm_first=True)
self.head = nn.TransformerEncoder(layer, head_layers, enable_nested_tensor=False) if head_layers else None
self.type_emb = nn.Embedding(3, width)
self.scorer = nn.Sequential(nn.LayerNorm(width), nn.Linear(width, width), nn.GELU(), nn.Linear(width, 1))
self.act_head = nn.Sequential(nn.Linear(width + 4, 256), nn.GELU(), nn.Linear(256, n_act))
self.register_buffer('temperature', torch.ones(3))
self.head_checkpointing = False
self.marker_only_head = False
self.encoder.config.reference_compile = False
@classmethod
def from_backbone(cls, path, revision=None, head_layers=2):
encoder = AutoModel.from_pretrained(path, revision=revision, trust_remote_code=False,
attn_implementation='sdpa')
return cls(encoder, head_layers=head_layers)
def forward(self, input_ids, attention_mask, marker_pos, marker_mask, qtype, return_actions=False):
hidden = self.encoder(input_ids=input_ids, attention_mask=attention_mask).last_hidden_state
hidden = hidden + self.type_emb(qtype)[:, None, :]
# Only option positions (and CLS for the action head) are consumed.
# The final block can keep full K/V context while avoiding unused queries
# and feed-forward outputs. Training retains the original dropout path.
selected = marker_pos
if return_actions:
selected = torch.cat((torch.zeros_like(marker_pos[:, :1]), marker_pos), dim=1)
sparse = (self.marker_only_head and not self.training and self.head is not None
and len(self.head.layers) > 0
and type(self.head.layers[-1]) is nn.TransformerEncoderLayer)
if self.head is not None:
padding = ~attention_mask.bool()
for i, layer in enumerate(self.head.layers):
if sparse and i == len(self.head.layers) - 1:
hidden = self._selected_head(layer, hidden, selected, attention_mask)
elif self.head_checkpointing and self.training:
hidden = checkpoint(layer, hidden, src_key_padding_mask=padding, use_reentrant=False)
else:
hidden = layer(hidden, src_key_padding_mask=padding)
if not sparse:
positions = selected[:, :, None].expand(-1, -1, hidden.shape[-1])
hidden = hidden.gather(1, positions)
markers = hidden[:, 1:] if return_actions else hidden
scores = self.scorer(markers).squeeze(-1).float()
scores = scores.masked_fill(~marker_mask, -1e4)
if not return_actions:
return scores
probability = scores.detach().softmax(-1)
top = probability.topk(min(2, probability.shape[1]), dim=-1).values
if top.shape[1] == 1:
top = torch.cat((top, torch.zeros_like(top)), dim=-1)
count = marker_mask.sum(-1).clamp_min(2).float()
entropy = -(probability * probability.clamp_min(1e-9).log()).sum(-1) / count.log()
features = torch.stack((top[:, 0], top[:, 0] - top[:, 1], entropy, count / 255), -1)
actions = self.act_head(torch.cat((hidden[:, 0].float(), features), -1))
return scores, actions
@staticmethod
def _selected_head(layer, hidden, selected, attention_mask):
"""Exact pre-norm block restricted to selected output positions (eval only)."""
width = hidden.shape[-1]
positions = selected[:, :, None].expand(-1, -1, width)
normalized = layer.norm1(hidden)
attn = layer.self_attn
query = F.linear(normalized.gather(1, positions),
attn.in_proj_weight[:width], attn.in_proj_bias[:width])
key, value = F.linear(normalized, attn.in_proj_weight[width:],
attn.in_proj_bias[width:]).chunk(2, dim=-1)
batch = hidden.shape[0]
heads = attn.num_heads
split = lambda x: x.reshape(batch, -1, heads, width // heads).transpose(1, 2)
attended = F.scaled_dot_product_attention(
split(query), split(key), split(value),
attn_mask=attention_mask[:, None, None, :].bool(), dropout_p=0.)
attended = attended.transpose(1, 2).reshape(batch, selected.shape[1], width)
result = hidden.gather(1, positions) + attn.out_proj(attended)
return result + layer.linear2(layer.activation(layer.linear1(layer.norm2(result))))
def save_pretrained(self, directory):
root = Path(directory)
root.mkdir(parents=True, exist_ok=True)
self.encoder.config.save_pretrained(root / 'encoder')
config = dict(format_version=1, architecture='JuliaDecisionModel',
weight_dtype=str(next(self.parameters()).dtype).replace('torch.', ''), **self.settings)
(root / 'julia_config.json').write_text(json.dumps(config, indent=2) + '\n')
save_file({k: v.detach().cpu().contiguous() for k, v in self.state_dict().items()},
str(root / 'model.safetensors'), metadata={'format': 'pt', 'family': 'julia'})
@classmethod
def from_pretrained(cls, directory, *, memory_map=False):
root = Path(directory)
if (root / 'julia_config.json').exists():
config = json.loads((root / 'julia_config.json').read_text())
if config['format_version'] != 1:
raise ValueError('Unsupported Julia checkpoint format')
settings = {k: config[k] for k in ('head_layers', 'n_act', 'dropout')}
else:
config = json.loads((root / 'rl_agent_config.json').read_text())
settings = dict(head_layers=config['head_layers'], n_act=len(config.get('act_costs', {})) + 1)
encoder_config = AutoConfig.from_pretrained(root / 'encoder', trust_remote_code=False)
encoder_config.reference_compile = False
try:
from transformers.initialization import no_init_weights
except ImportError:
from transformers.modeling_utils import no_init_weights
with no_init_weights():
encoder = AutoModel.from_config(encoder_config, attn_implementation='sdpa', trust_remote_code=False)
model = cls(encoder, **settings)
if not memory_map and config.get('weight_dtype') == 'bfloat16':
model.to(dtype=torch.bfloat16)
# Inference can retain safetensors' private file-backed storage rather than
# copying the entire (mostly unused vocabulary) embedding into anonymous RAM.
# The default copy path remains available for training and mutable checkpoints.
model.load_state_dict(load_file(str(root / 'model.safetensors')),
strict=True, assign=memory_map)
return model