Qyvos / julia /router /encoder.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
1.79 kB
"""Inference-only ModernBERT path for decision models (no unused outputs)."""
from types import MethodType
import torch
from transformers.modeling_outputs import BaseModelOutput
def _decision_forward(self, input_ids=None, attention_mask=None, **kwargs):
if self.training or kwargs or input_ids is None or attention_mask is None:
return self._julia_original_forward(input_ids=input_ids, attention_mask=attention_mask, **kwargs)
position_ids = torch.arange(input_ids.shape[1], device=input_ids.device).unsqueeze(0)
full_mask, local_mask = self._update_attention_mask(attention_mask, output_attentions=False)
hidden = self.embeddings(input_ids=input_ids)
# Upstream iterates config.layer_types (one entry per layer), overwriting
# the same two dictionary entries 22 times in this checkpoint.
positions = {kind: self.rotary_emb(hidden, position_ids, kind) for kind in self._julia_attention_types}
for layer in self.layers:
hidden = layer(hidden, attention_mask=full_mask, sliding_window_mask=local_mask,
position_ids=position_ids, cu_seqlens=None, max_seqlen=None,
position_embeddings=positions[layer.attention_type], output_attentions=False)[0]
return BaseModelOutput(last_hidden_state=self.final_norm(hidden))
def specialize_decision_encoder(model):
encoder = model.encoder
if encoder.config.model_type != 'modernbert' or encoder.config._attn_implementation != 'sdpa':
return False
if hasattr(encoder, '_julia_original_forward'):
return True
encoder._julia_original_forward = encoder.forward
encoder._julia_attention_types = tuple(dict.fromkeys(encoder.config.layer_types))
encoder.forward = MethodType(_decision_forward, encoder)
return True