Text Classification
Transformers
Safetensors
Thai
English
openthai_systemone
feature-extraction
system-one
decision-model
thai
qwen3.5
quantized
compressed-tensors
llm-compressor
custom_code
Instructions to use iapp/OpenThai-SystemOne-FP8-Dynamic with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use iapp/OpenThai-SystemOne-FP8-Dynamic with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("text-classification", model="iapp/OpenThai-SystemOne-FP8-Dynamic", trust_remote_code=True)# Load model directly from transformers import AutoModel model = AutoModel.from_pretrained("iapp/OpenThai-SystemOne-FP8-Dynamic", trust_remote_code=True, device_map="auto") - Notebooks
- Google Colab
- Kaggle
| """OpenThai-SystemOne decision model. | |
| text tower (Qwen3.5-0.8B, LM head removed) -> hidden state at every <|ts_answer|> -> SlotHead (256 logits) | |
| mask slots >= k -> softmax -> probabilities over the k options | |
| -Vision variant: the tower is the multimodal Qwen3.5 model (ViT + merger + the same text tower). Images enter as | |
| runs of <|image_pad|> tokens; `point` questions are read out by the PointHead, which attends from the <|ts_answer|> | |
| hidden state over the LLM's final hidden states of that image's visual tokens (GUI-Actor recipe) plus a learned | |
| "null" key for "not on the image". With no pixel_values the forward path is identical to the text-only model. | |
| """ | |
| from __future__ import annotations | |
| import math | |
| from dataclasses import dataclass | |
| from typing import Optional | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| from transformers import AutoModel, AutoModelForCausalLM, AutoTokenizer, PreTrainedModel | |
| from transformers.utils import ModelOutput | |
| from .configuration import OpenThaiSystemOneConfig | |
| from .formatting import SPECIAL_TOKENS, TOK_ANSWER, TOK_POINT, QWEN_IMAGE_PAD, QWEN_VISION_END, QWEN_VISION_START, add_special_tokens | |
| QTYPE_INDEX = {"choice": 0, "score": 1, "noul": 2, "point": 3} | |
| class DecisionOutput(ModelOutput): | |
| loss: Optional[torch.Tensor] = None | |
| logits: Optional[torch.Tensor] = None # (B, Q, n_slots), masked with -inf | |
| probs: Optional[torch.Tensor] = None # (B, Q, n_slots) | |
| hidden_states: Optional[torch.Tensor] = None # (B, Q, H) at answer positions | |
| point_logits: Optional[torch.Tensor] = None # (B, Q, L+1): per visual token of the referenced image + null; -inf elsewhere | |
| point_probs: Optional[torch.Tensor] = None | |
| slot_loss: Optional[torch.Tensor] = None | |
| point_loss: Optional[torch.Tensor] = None | |
| class PointHead(nn.Module): | |
| """Attention readout over one image's visual tokens (final-layer LLM states). | |
| q = MLP_T(h_answer), k_i = MLP_V(SA(v_i)), logits_i = q.k_i / sqrt(d) / T ; extra null key = "not on the image". | |
| """ | |
| def __init__(self, hidden: int, dim: int = 256, n_sa_layers: int = 1, n_heads: int = 8): | |
| super().__init__() | |
| self.sa = nn.ModuleList( | |
| [nn.TransformerEncoderLayer(hidden, n_heads, dim_feedforward=2 * hidden, dropout=0.0, batch_first=True, norm_first=True) for _ in range(n_sa_layers)] | |
| ) | |
| self.q_proj = nn.Sequential(nn.Linear(hidden, dim), nn.GELU(), nn.Linear(dim, dim)) | |
| self.k_proj = nn.Sequential(nn.Linear(hidden, dim), nn.GELU(), nn.Linear(dim, dim)) | |
| self.null_key = nn.Parameter(torch.zeros(dim)) | |
| self.log_temperature = nn.Parameter(torch.zeros(())) | |
| self.dim = dim | |
| def forward(self, h_answer: torch.Tensor, visual: torch.Tensor, visual_mask: torch.Tensor) -> torch.Tensor: | |
| """h_answer (N, H); visual (N, L, H) right-padded; visual_mask (N, L) True = real token. Returns (N, L+1) logits.""" | |
| x = visual | |
| for layer in self.sa: | |
| x = layer(x, src_key_padding_mask=~visual_mask) | |
| k = self.k_proj(x) # (N, L, d) | |
| q = self.q_proj(h_answer) # (N, d) | |
| logits = torch.einsum("nd,nld->nl", q, k) / math.sqrt(self.dim) | |
| null = (q @ self.null_key)[:, None] / math.sqrt(self.dim) | |
| logits = torch.cat([logits, null], dim=1) / self.log_temperature.exp() | |
| pad = torch.cat([~visual_mask, torch.zeros_like(visual_mask[:, :1])], dim=1) | |
| return logits.masked_fill(pad, float("-inf")) | |
| class OpenThaiSystemOneForDecision(PreTrainedModel): | |
| config_class = OpenThaiSystemOneConfig | |
| base_model_prefix = "model" | |
| supports_gradient_checkpointing = True | |
| _supports_flash_attn = True | |
| _supports_sdpa = True | |
| def __init__(self, config: OpenThaiSystemOneConfig): | |
| super().__init__(config) | |
| # Parameter layout is the SAME as the text-only line for everything they share: text tower `model.*`, | |
| # `slot_head.*`, 3 `log_temperature`s. The vision parts sit beside it (`visual.*`, `point_head.*`), so a | |
| # text-only v0.x client loading a -Vision checkpoint finds every weight it knows and ignores the rest. | |
| self.model = AutoModel.from_config(config.text_config) | |
| if config.is_vision: | |
| self.visual = AutoModel.from_config(config.vision_config) | |
| enable_vision_helper_caches() | |
| ph = dict(config.point_head or {}) | |
| self.point_head = PointHead(config.hidden_size, ph.get("dim", 256), ph.get("n_sa_layers", 1), ph.get("n_heads", 8)) | |
| else: | |
| self.visual = None | |
| self.point_head = None | |
| self.slot_head = nn.Linear(config.hidden_size, config.n_slots, bias=config.head_bias) | |
| # log-temperatures per question type (choice/score/noul[/point]); learned in the calibration stage | |
| self.log_temperature = nn.Parameter(torch.zeros(config.n_temperatures)) | |
| self.post_init() | |
| def language_model(self): | |
| return self.model | |
| def encode(self, input_ids, attention_mask=None, pixel_values=None, image_grid_thw=None, position_ids=None, **kwargs): | |
| """Final hidden states. With images: ViT -> merger -> splice into the <|image_pad|> positions, 3-D (t,h,w) position | |
| ids (computed here when not given), then the text tower. Without images: the plain text-tower forward.""" | |
| if pixel_values is None: | |
| return self.model(input_ids=input_ids, attention_mask=attention_mask, position_ids=position_ids, **kwargs).last_hidden_state | |
| if self.visual is None: | |
| raise ValueError("this checkpoint is text-only; images need OpenThai-SystemOne-Vision") | |
| _VISION_TARGET_DEVICE[0] = pixel_values.device | |
| embeds = self.model.get_input_embeddings()(input_ids) | |
| img = self.visual(pixel_values.to(self.visual.dtype), grid_thw=image_grid_thw).pooler_output | |
| mask = (input_ids == self.config.image_token_id).unsqueeze(-1) | |
| if int(mask.sum()) != img.shape[0]: | |
| raise ValueError(f"image tokens ({int(mask.sum())}) != image features ({img.shape[0]})") | |
| embeds = embeds.masked_scatter(mask.to(embeds.device), img.to(embeds.device, embeds.dtype)) | |
| if position_ids is None: | |
| position_ids = mrope_positions_from_ids(input_ids, attention_mask, image_grid_thw, self.config.image_token_id, | |
| int(getattr(self.config.vision_config, "spatial_merge_size", 2))) | |
| return self.model(inputs_embeds=embeds, attention_mask=attention_mask, position_ids=position_ids.to(embeds.device), **kwargs).last_hidden_state | |
| def get_input_embeddings(self): | |
| return self.language_model.get_input_embeddings() | |
| def set_input_embeddings(self, value): | |
| self.language_model.set_input_embeddings(value) | |
| # ------------------------------------------------------------------ construction | |
| def from_causal_lm( | |
| cls, | |
| path: str, | |
| *, | |
| tokenizer=None, | |
| n_slots: int = 256, | |
| torch_dtype=torch.bfloat16, | |
| **kwargs, | |
| ): | |
| """Build a text-only decision model from a (text-only) causal-LM checkpoint: drop lm_head, add tokens + head.""" | |
| tok = tokenizer or AutoTokenizer.from_pretrained(path) | |
| added = add_special_tokens(tok) | |
| lm = AutoModelForCausalLM.from_pretrained(path, dtype=torch_dtype, **kwargs) | |
| base = lm.model if hasattr(lm, "model") else lm.base_model | |
| text_cfg = base.config | |
| if added: | |
| lm.resize_token_embeddings(len(tok), mean_resizing=False) | |
| text_cfg.vocab_size = lm.get_input_embeddings().weight.shape[0] | |
| _init_new_token_embeddings(lm.get_input_embeddings().weight, tok, added) | |
| cfg = OpenThaiSystemOneConfig( | |
| text_config=text_cfg, | |
| n_slots=n_slots, | |
| answer_token_id=tok.convert_tokens_to_ids(TOK_ANSWER), | |
| pad_token_id=tok.pad_token_id, | |
| ) | |
| cfg.text_config.tie_word_embeddings = False # there is no LM head any more | |
| model = cls(cfg).to(torch_dtype) | |
| missing, unexpected = model.model.load_state_dict(base.state_dict(), strict=False) | |
| assert not unexpected, unexpected | |
| _init_slot_head(model.slot_head) | |
| model.model.config = cfg.text_config | |
| return model, tok | |
| def from_text_decision_model( | |
| cls, | |
| text_path: str, | |
| vision_source: str, | |
| *, | |
| tokenizer=None, | |
| point_head: Optional[dict] = None, | |
| torch_dtype=torch.bfloat16, | |
| ): | |
| """Build the -Vision model: text tower + slot head from a text-only decision checkpoint (e.g. v0.3), the ViT + | |
| merger from a Qwen3.5 multimodal checkpoint (base/Qwen3.5-0.8B-Base), a fresh PointHead, and the two extra | |
| control tokens. Text-only forward of the result equals the source model.""" | |
| import glob | |
| import json | |
| import os | |
| from safetensors import safe_open | |
| from transformers import AutoConfig | |
| text_model = cls.from_pretrained(text_path, dtype=torch_dtype) | |
| assert not text_model.config.is_vision, "source must be the text-only model" | |
| tok = tokenizer or AutoTokenizer.from_pretrained(text_path) | |
| added = add_special_tokens(tok, vision=True) | |
| src_cfg = AutoConfig.from_pretrained(vision_source) | |
| vcfg = src_cfg.vision_config | |
| vcfg.out_hidden_size = text_model.config.hidden_size | |
| cfg = OpenThaiSystemOneConfig( | |
| text_config=text_model.config.text_config, | |
| n_slots=text_model.config.n_slots, | |
| abstain_slot=text_model.config.abstain_slot, | |
| answer_token_id=text_model.config.answer_token_id, | |
| head_bias=text_model.config.head_bias, | |
| n_temperatures=text_model.config.n_temperatures, | |
| vision_config=vcfg, | |
| image_token_id=tok.convert_tokens_to_ids(QWEN_IMAGE_PAD), | |
| vision_start_token_id=tok.convert_tokens_to_ids(QWEN_VISION_START), | |
| vision_end_token_id=tok.convert_tokens_to_ids(QWEN_VISION_END), | |
| point_token_id=tok.convert_tokens_to_ids(TOK_POINT), | |
| point_head=point_head or {"dim": 256, "n_sa_layers": 1, "n_heads": 8}, | |
| pad_token_id=tok.pad_token_id, | |
| ) | |
| cfg.text_config.vocab_size = len(tok) | |
| model = cls(cfg).to(torch_dtype) | |
| # text tower + heads | |
| old_vocab = text_model.get_input_embeddings().weight.shape[0] | |
| if len(tok) > old_vocab: | |
| text_model.model.resize_token_embeddings(len(tok), mean_resizing=False) | |
| _init_new_token_embeddings(text_model.get_input_embeddings().weight, tok, len(tok) - old_vocab) | |
| missing, unexpected = model.model.load_state_dict(text_model.model.state_dict(), strict=False) | |
| assert not unexpected and not missing, (missing, unexpected) | |
| model.slot_head.load_state_dict(text_model.slot_head.state_dict()) | |
| with torch.no_grad(): | |
| n = min(text_model.log_temperature.numel(), model.log_temperature.numel()) | |
| model.log_temperature[:n] = text_model.log_temperature[:n] | |
| # vision tower from the multimodal checkpoint | |
| files = sorted(glob.glob(os.path.join(vision_source, "*.safetensors"))) | |
| sd = {} | |
| for f in files: | |
| with safe_open(f, "pt") as fh: | |
| for k in fh.keys(): | |
| for prefix in ("model.visual.", "visual."): | |
| if k.startswith(prefix): | |
| sd[k[len(prefix):]] = fh.get_tensor(k) | |
| break | |
| missing, unexpected = model.visual.load_state_dict({k: v.to(torch_dtype) for k, v in sd.items()}, strict=False) | |
| assert not unexpected and not missing, (missing, unexpected) | |
| model.model.config = cfg.text_config | |
| return model, tok | |
| # ------------------------------------------------------------------ forward | |
| def gather_answer_states(self, hidden: torch.Tensor, answer_positions: torch.Tensor) -> torch.Tensor: | |
| idx = answer_positions.clamp(min=0).unsqueeze(-1).expand(-1, -1, hidden.shape[-1]) | |
| return torch.gather(hidden, 1, idx) # (B, Q, H) | |
| def slot_logits( | |
| self, | |
| answer_hidden: torch.Tensor, | |
| option_counts: torch.Tensor, | |
| *, | |
| include_abstain: bool = True, | |
| qtypes: Optional[torch.Tensor] = None, | |
| apply_temperature: bool = True, | |
| ) -> torch.Tensor: | |
| logits = self.slot_head(answer_hidden.to(self.slot_head.weight.dtype)).float() | |
| if apply_temperature: | |
| if qtypes is None: | |
| t = self.log_temperature[0].exp() | |
| else: | |
| t = self.log_temperature.exp()[qtypes.clamp(min=0, max=self.log_temperature.numel() - 1)] # (B, Q) | |
| t = t.unsqueeze(-1) | |
| logits = logits / t | |
| ar = torch.arange(logits.shape[-1], device=logits.device) | |
| valid = ar[None, None, :] < option_counts.unsqueeze(-1) | |
| if include_abstain: | |
| valid = valid.clone() | |
| valid[..., self.config.abstain_slot] = True | |
| # questions that are padding or point questions (option_counts == 0) keep slot 0 valid to avoid NaNs | |
| valid[..., 0] |= option_counts.unsqueeze(-1).squeeze(-1) == 0 | |
| return logits.masked_fill(~valid, float("-inf")) | |
| def point_logits( | |
| self, | |
| hidden: torch.Tensor, | |
| answer_hidden: torch.Tensor, | |
| image_spans: torch.Tensor, | |
| point_image_index: torch.Tensor, | |
| *, | |
| max_tokens: Optional[int] = None, | |
| ) -> Optional[torch.Tensor]: | |
| """(B, Q, L+1) logits for point questions (-inf rows elsewhere).""" | |
| B, Q = point_image_index.shape | |
| sel = (point_image_index >= 0).nonzero(as_tuple=False) | |
| L = max_tokens or int((image_spans[..., 1] - image_spans[..., 0]).clamp(min=0).max().item()) | |
| out = torch.full((B, Q, L + 1), float("-inf"), device=hidden.device) | |
| if sel.numel() == 0: | |
| return out | |
| n = sel.shape[0] | |
| vis = hidden.new_zeros((n, L, hidden.shape[-1])) | |
| mask = torch.zeros((n, L), dtype=torch.bool, device=hidden.device) | |
| for r, (b, q) in enumerate(sel.tolist()): | |
| s0, e0 = image_spans[b, point_image_index[b, q]].tolist() | |
| vis[r, : e0 - s0] = hidden[b, s0:e0] | |
| mask[r, : e0 - s0] = True | |
| h = answer_hidden[sel[:, 0], sel[:, 1]] | |
| pl = self.point_head(h.to(vis.dtype), vis, mask).float() | |
| out[sel[:, 0], sel[:, 1]] = pl | |
| return out | |
| def forward( | |
| self, | |
| input_ids: torch.Tensor, | |
| attention_mask: Optional[torch.Tensor] = None, | |
| answer_positions: Optional[torch.Tensor] = None, | |
| option_counts: Optional[torch.Tensor] = None, | |
| labels: Optional[torch.Tensor] = None, | |
| soft_labels: Optional[torch.Tensor] = None, | |
| qtypes: Optional[torch.Tensor] = None, | |
| include_abstain: bool = True, | |
| label_smoothing: float = 0.0, | |
| brier_weight: float = 0.0, | |
| apply_temperature: bool = True, | |
| pixel_values: Optional[torch.Tensor] = None, | |
| image_grid_thw: Optional[torch.Tensor] = None, | |
| mm_token_type_ids: Optional[torch.Tensor] = None, | |
| image_spans: Optional[torch.Tensor] = None, | |
| point_image_index: Optional[torch.Tensor] = None, | |
| point_targets: Optional[torch.Tensor] = None, | |
| point_has_target: Optional[torch.Tensor] = None, | |
| point_loss_weight: float = 1.0, | |
| position_ids: Optional[torch.Tensor] = None, | |
| **kwargs, | |
| ) -> DecisionOutput: | |
| # position_ids (3, B, T) from formatting.mrope_position_ids avoid recomputing them; image_grid_thw may stay on the CPU | |
| hidden = self.encode(input_ids, attention_mask=attention_mask, pixel_values=pixel_values, image_grid_thw=image_grid_thw, | |
| position_ids=position_ids if pixel_values is not None else None, **kwargs) | |
| if answer_positions is None: | |
| answer_positions = (input_ids == self.config.answer_token_id).nonzero()[:, 1].unsqueeze(0) | |
| if option_counts is None: | |
| raise ValueError("option_counts required") | |
| h = self.gather_answer_states(hidden, answer_positions) | |
| logits = self.slot_logits(h, option_counts, include_abstain=include_abstain, qtypes=qtypes, apply_temperature=apply_temperature) | |
| probs = logits.softmax(-1) | |
| p_logits = p_probs = None | |
| if self.point_head is not None and point_image_index is not None and (point_image_index >= 0).any(): | |
| p_logits = self.point_logits(hidden, h, image_spans, point_image_index, | |
| max_tokens=(point_targets.shape[-1] - 1) if point_targets is not None else None) | |
| p_probs = p_logits.softmax(-1) | |
| p_probs = p_probs.masked_fill(torch.isinf(p_logits).all(-1, keepdim=True), 0.0) | |
| loss = slot_loss = point_loss = None | |
| if labels is not None or soft_labels is not None: | |
| logp = logits.log_softmax(-1) | |
| if soft_labels is not None: | |
| valid = (option_counts > 0) | |
| tgt = soft_labels.float() | |
| nll = -(tgt * logp.masked_fill(torch.isinf(logp), 0.0)).sum(-1) | |
| slot_loss = (nll * valid).sum() / valid.sum().clamp(min=1) | |
| else: | |
| flat_logp = logp.reshape(-1, logp.shape[-1]) | |
| flat_lab = labels.reshape(-1) | |
| keep = flat_lab != -100 | |
| if keep.any(): | |
| lp = flat_logp[keep] | |
| lb = flat_lab[keep] | |
| nll = -lp.gather(1, lb[:, None]).squeeze(1) | |
| if label_smoothing > 0: | |
| n_valid = torch.isfinite(lp).sum(-1).clamp(min=1).float() | |
| smooth = -(lp.masked_fill(torch.isinf(lp), 0.0)).sum(-1) / n_valid | |
| nll = (1 - label_smoothing) * nll + label_smoothing * smooth | |
| slot_loss = nll.mean() | |
| if brier_weight > 0: | |
| p = lp.exp() | |
| onehot = F.one_hot(lb, p.shape[-1]).float() | |
| slot_loss = slot_loss + brier_weight * ((p - onehot) ** 2).sum(-1).mean() | |
| else: | |
| slot_loss = logits.sum() * 0.0 | |
| loss = slot_loss | |
| if p_logits is not None and point_targets is not None and point_has_target is not None and point_has_target.any(): | |
| lp = p_logits.log_softmax(-1).masked_fill(torch.isinf(p_logits), 0.0) | |
| tgt = point_targets.float() | |
| nll = -(tgt * lp).sum(-1) # cross-entropy against the soft mask (= KL up to the target entropy) | |
| point_loss = (nll * point_has_target).sum() / point_has_target.sum().clamp(min=1) | |
| loss = point_loss * point_loss_weight if loss is None else loss + point_loss_weight * point_loss | |
| return DecisionOutput(loss=loss, logits=logits, probs=probs, hidden_states=h, point_logits=p_logits, point_probs=p_probs, | |
| slot_loss=slot_loss, point_loss=point_loss) | |
| _VISION_CACHE_ON = False | |
| _VISION_TARGET_DEVICE = [None] # set by the forward: cached ViT helper outputs are moved here (lets image_grid_thw stay on the CPU) | |
| def enable_vision_helper_caches(maxsize: int = 512): | |
| """Memoise the Qwen3.5 ViT's per-call helpers by image grid. | |
| `get_vision_interpolation_indices_and_weights`, `get_vision_position_ids` and `get_vision_attention_seqlens` rebuild | |
| the same index tensors on every forward (tens of ms of small ops for a 320x240 frame). Their outputs depend only on | |
| the grid (t, h, w) and the device, so they are cached; a Doom loop or a fixed-size screenshot stream hits the cache | |
| every step. Idempotent; called automatically when a vision model is built. | |
| """ | |
| global _VISION_CACHE_ON | |
| if _VISION_CACHE_ON: | |
| return | |
| try: | |
| from transformers.models.qwen3_5 import modeling_qwen3_5 as mq | |
| except Exception: # pragma: no cover | |
| return | |
| def cached(fn, key_extra=lambda *a, **k: ()): | |
| cache = {} | |
| def wrapper(grid_thw, *args, **kwargs): | |
| kw = {k: v for k, v in kwargs.items() if k != "kwargs"} | |
| target = _VISION_TARGET_DEVICE[0] or grid_thw.device | |
| key = (tuple(map(tuple, grid_thw.tolist())), str(target), tuple(sorted((k, str(v)) for k, v in kw.items())), tuple(str(a) for a in args)) | |
| hit = cache.get(key) | |
| if hit is None: | |
| hit = fn(grid_thw, *args, **kwargs) | |
| def mv(x): | |
| return x.to(target) if torch.is_tensor(x) else x | |
| hit = tuple(mv(x) for x in hit) if isinstance(hit, tuple) else mv(hit) | |
| if len(cache) >= maxsize: | |
| cache.clear() | |
| cache[key] = hit | |
| return hit | |
| wrapper.__wrapped__ = fn | |
| return wrapper | |
| for name in ("get_vision_interpolation_indices_and_weights", "get_vision_position_ids", "get_vision_attention_seqlens"): | |
| fn = getattr(mq, name, None) | |
| if fn is not None and not hasattr(fn, "__wrapped__"): | |
| setattr(mq, name, cached(fn)) | |
| _VISION_CACHE_ON = True | |
| def mrope_positions_from_ids(input_ids, attention_mask, image_grid_thw, image_token_id: int, merge: int): | |
| """3-D position ids from token ids alone (same result as formatting.mrope_position_ids / Qwen's get_rope_index).""" | |
| B, T = input_ids.shape | |
| grids = [[int(v) for v in g] for g in image_grid_thw.tolist()] | |
| ids = input_ids.tolist() | |
| am = attention_mask.tolist() if attention_mask is not None else [[1] * T for _ in range(B)] | |
| pos = torch.zeros((3, B, T), dtype=torch.long) | |
| gi = 0 | |
| for b in range(B): | |
| cur, i = 0, 0 | |
| n = sum(am[b]) | |
| while i < n: | |
| if ids[b][i] == image_token_id: | |
| t, h, w = grids[gi]; gi += 1 | |
| hm, wm = h // merge, w // merge | |
| L = t * hm * wm | |
| tt = torch.arange(t).view(t, 1, 1).expand(t, hm, wm).reshape(-1) | |
| hh = torch.arange(hm).view(1, hm, 1).expand(t, hm, wm).reshape(-1) | |
| ww = torch.arange(wm).view(1, 1, wm).expand(t, hm, wm).reshape(-1) | |
| pos[0, b, i:i + L] = tt + cur; pos[1, b, i:i + L] = hh + cur; pos[2, b, i:i + L] = ww + cur | |
| cur += max(hm, wm); i += L | |
| else: | |
| j = i | |
| while j < n and ids[b][j] != image_token_id: | |
| j += 1 | |
| pos[:, b, i:j] = torch.arange(j - i) + cur | |
| cur += j - i; i = j | |
| return pos | |
| def convert_legacy_vision_state_dict(sd: dict) -> dict: | |
| """Checkpoints trained before 2026-09-23 stored the vision model as a Qwen3_5Model (`model.language_model.*`, | |
| `model.visual.*`, 4 temperatures). Map them onto the shared layout (`model.*`, `visual.*`, 3 temperatures).""" | |
| out = {} | |
| for k, v in sd.items(): | |
| if k.startswith("model.language_model."): | |
| out["model." + k[len("model.language_model."):]] = v | |
| elif k.startswith("model.visual."): | |
| out["visual." + k[len("model.visual."):]] = v | |
| elif k == "log_temperature" and v.numel() == 4: | |
| out[k] = v[:3].clone() | |
| else: | |
| out[k] = v | |
| return out | |
| def use_reference_kernels(): | |
| """Force the pure-PyTorch Gated-DeltaNet / causal-conv paths. | |
| transformers routes `chunk_gated_delta_rule` & co. to the Triton kernels (flash-linear-attention, causal-conv1d) | |
| whenever those packages are importable, without checking the tensor device, which crashes on CPU/MPS. | |
| Call this before running on a non-CUDA device. | |
| """ | |
| try: | |
| from transformers.models.qwen3_5 import modeling_qwen3_5 as m | |
| except Exception: # pragma: no cover | |
| return | |
| for name in ("torch_chunk_gated_delta_rule", "torch_recurrent_gated_delta_rule", "chunk_gated_delta_rule", | |
| "fused_recurrent_gated_delta_rule", "causal_conv1d_fn", "causal_conv1d_update"): | |
| fn = getattr(m, name, None) | |
| if fn is not None and hasattr(fn, "__wrapped__"): | |
| setattr(m, name, fn.__wrapped__) | |
| def _init_slot_head(head: nn.Linear): | |
| nn.init.normal_(head.weight, std=0.02) | |
| if head.bias is not None: | |
| nn.init.zeros_(head.bias) | |
| def _init_new_token_embeddings(weight: torch.Tensor, tok, n_added: int): | |
| """New control tokens start near the mean of digit-token embeddings + small noise.""" | |
| digit_ids = [tok.convert_tokens_to_ids(d) for d in "0123456789"] | |
| digit_ids = [i for i in digit_ids if i is not None and i != tok.unk_token_id] | |
| mean = weight[digit_ids].float().mean(0) if digit_ids else weight[: weight.shape[0] - n_added].float().mean(0) | |
| std = weight[: weight.shape[0] - n_added].float().std() | |
| new = mean[None, :] + torch.randn(n_added, weight.shape[1]) * std * 0.1 | |
| weight[-n_added:] = new.to(weight.dtype) | |
| def confidence_from_probs(p: torch.Tensor, k: int) -> float: | |
| """1 - normalised entropy over the k valid options.""" | |
| if k <= 1: | |
| return 1.0 | |
| p = p[:k].clamp(min=1e-12) | |
| p = p / p.sum() | |
| h = -(p * p.log()).sum().item() | |
| return float(max(0.0, min(1.0, 1.0 - h / math.log(k)))) | |