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
| """Turn (state, questions) into token ids + slot bookkeeping. | |
| Layout (one sequence, causal): | |
| <|ts_state|> {state text} | |
| <|ts_q|><|ts_choice|> {instructions} | |
| <|ts_opt_0|> {option name}: {description} | |
| <|ts_opt_1|> {option name} | |
| ... | |
| <|ts_answer|> <- hidden state here -> SlotHead (256 logits) | |
| <|ts_q|><|ts_noul|> {instructions} | |
| <|ts_opt_0|> no | |
| <|ts_opt_1|> yes | |
| <|ts_answer|> | |
| ... | |
| Slot i (0..254) means "the option introduced by <|ts_opt_i|>"; slot 255 = abstain. | |
| All answers for all questions are read out from one forward pass. | |
| Vision variant (OpenThai-SystemOne-Vision) adds, before the state (or inline where the state text says <image:id>): | |
| <|ts_img|> [image:screen 1280x800 screenshot] | |
| <|vision_start|><|image_pad|> x N <|vision_end|> | |
| and a fourth question type whose readout is the PointHead over that image's N visual tokens: | |
| <|ts_q|><|ts_point|> [image:screen] {instructions} | |
| <|ts_answer|> | |
| """ | |
| from __future__ import annotations | |
| import json | |
| import random | |
| from dataclasses import dataclass, field | |
| from typing import Any, Dict, List, Optional, Sequence, Tuple, Union | |
| from .types import Choice, Noul, Point, Score, Question, MAX_OPTIONS | |
| N_SLOTS = 256 | |
| ABSTAIN_SLOT = 255 | |
| TOK_STATE = "<|ts_state|>" | |
| TOK_Q = "<|ts_q|>" | |
| TOK_CHOICE = "<|ts_choice|>" | |
| TOK_SCORE = "<|ts_score|>" | |
| TOK_NOUL = "<|ts_noul|>" | |
| TOK_ANSWER = "<|ts_answer|>" | |
| TOK_OPT = [f"<|ts_opt_{i}|>" for i in range(N_SLOTS)] | |
| TOK_POINT = "<|ts_point|>" # vision variant only | |
| TOK_IMG = "<|ts_img|>" # vision variant only | |
| SPECIAL_TOKENS: List[str] = [TOK_STATE, TOK_Q, TOK_CHOICE, TOK_SCORE, TOK_NOUL, TOK_ANSWER] + TOK_OPT | |
| VISION_TOKENS: List[str] = [TOK_POINT, TOK_IMG] # added on top of SPECIAL_TOKENS for the -Vision variant | |
| # Qwen's own multimodal tokens (already in the base tokenizer) | |
| QWEN_VISION_START, QWEN_VISION_END, QWEN_IMAGE_PAD = "<|vision_start|>", "<|vision_end|>", "<|image_pad|>" | |
| QTYPES = ("choice", "score", "noul", "point") | |
| NOUL_OPTIONS = ("no", "yes") # slot 0 = no, slot 1 = yes -> noul = p(slot 1) | |
| DEFAULT_MAX_TOTAL_TOKENS = 65536 | |
| DEFAULT_MAX_STATE_TOKENS = 32768 | |
| def add_special_tokens(tokenizer, *, vision: bool = False) -> int: | |
| """Register the control tokens (plus the vision ones when vision=True). Returns number of tokens added.""" | |
| existing = set(tokenizer.get_vocab()) | |
| wanted = SPECIAL_TOKENS + (VISION_TOKENS if vision else []) | |
| new = [t for t in wanted if t not in existing] | |
| if not new: | |
| return 0 | |
| return tokenizer.add_tokens(new, special_tokens=True) | |
| def sanitize(text: str) -> str: | |
| """Stop user content from smuggling control tokens into the sequence.""" | |
| return text.replace("<|ts_", "<|ts_") if "<|ts_" in text else text | |
| def state_to_text(state: Union[str, Dict[str, Any], List[Any]], *, indent: Optional[int] = None) -> str: | |
| if isinstance(state, str): | |
| return state | |
| return json.dumps(state, ensure_ascii=False, indent=indent) | |
| class QuestionSpec: | |
| """A question flattened to option strings + slot bookkeeping.""" | |
| qid: str | |
| qtype: str # choice | score | noul | |
| instructions: str | |
| option_names: List[str] # in slot order (after any permutation) | |
| option_descs: List[Optional[str]] | |
| perm: List[int] # perm[slot] = original index of the option at that slot | |
| label_slot: Optional[int] = None # training only | |
| image: Optional[str] = None # point questions: id of the referenced image (None = first image) | |
| point_bbox: Optional[List[float]] = None # training only: normalised [x0,y0,x1,y1]; None with point_label_given -> null | |
| point_label_given: bool = False | |
| def question_to_spec( | |
| qid: str, | |
| q: Question, | |
| *, | |
| label: Optional[Union[str, int, bool]] = None, | |
| shuffle: bool = False, | |
| rng: Optional[random.Random] = None, | |
| drop_label: bool = False, | |
| perm: Optional[Sequence[int]] = None, | |
| ) -> QuestionSpec: | |
| """Flatten a typed question. | |
| label: for training. Choice -> option name; Score -> level index (int); Noul -> bool. | |
| shuffle: permute option order (Choice only; Score/Noul order is semantic). | |
| drop_label: remove the correct option from a Choice so the target becomes ABSTAIN_SLOT. | |
| perm: explicit option order for a Choice (list of original indices), e.g. a cyclic shift for order-invariant inference. | |
| """ | |
| if isinstance(q, Choice): | |
| names = list(q.criteria.keys()) | |
| descs = [q.criteria[n] for n in names] | |
| idx = list(range(len(names))) | |
| label_idx = None | |
| if label is not None: | |
| if label not in q.criteria: | |
| raise ValueError(f"label {label!r} is not one of the options") | |
| label_idx = names.index(str(label)) | |
| if drop_label and label_idx is not None: | |
| if len(idx) < 2: | |
| raise ValueError("cannot drop the only option") | |
| idx.remove(label_idx) | |
| label_idx = None | |
| if perm is not None: | |
| idx = [i for i in perm if i in idx] | |
| elif shuffle: | |
| (rng or random).shuffle(idx) | |
| names_p = [names[i] for i in idx] | |
| descs_p = [descs[i] for i in idx] | |
| if label is None: | |
| slot = None | |
| elif label_idx is None: | |
| slot = ABSTAIN_SLOT | |
| else: | |
| slot = idx.index(label_idx) | |
| return QuestionSpec(qid, "choice", q.instructions, names_p, descs_p, idx, slot) | |
| if isinstance(q, Score): | |
| names = [str(i) for i in range(len(q.criteria))] | |
| descs = list(q.criteria) | |
| slot = int(label) if label is not None else None | |
| if slot is not None and not (0 <= slot < len(descs)): | |
| raise ValueError("score label out of range") | |
| return QuestionSpec(qid, "score", q.instructions, names, descs, list(range(len(names))), slot) | |
| if isinstance(q, Noul): | |
| c = q.criteria or {} | |
| descs = [c.get("false"), c.get("true")] | |
| slot = None if label is None else int(bool(label)) | |
| return QuestionSpec(qid, "noul", q.instructions, list(NOUL_OPTIONS), descs, [0, 1], slot) | |
| if isinstance(q, Point): | |
| # label: {"bbox": [x0,y0,x1,y1]} | [x0,y0,x1,y1] | {"bbox": None} (= not on screen) | None (inference) | |
| bbox, given = None, False | |
| if label is not None: | |
| given = True | |
| bbox = label.get("bbox") if isinstance(label, dict) else list(label) | |
| if bbox is not None: | |
| bbox = [float(v) for v in bbox] | |
| if len(bbox) != 4: | |
| raise ValueError("point label bbox needs 4 numbers") | |
| return QuestionSpec(qid, "point", q.instructions, [], [], [], None, image=q.image, point_bbox=bbox, point_label_given=given) | |
| raise TypeError(type(q)) | |
| def spec_to_text(spec: QuestionSpec, *, image_id: Optional[str] = None) -> str: | |
| if spec.qtype == "point": | |
| tag = f"[image:{sanitize(str(image_id if image_id is not None else spec.image or ''))}] " | |
| return f"{TOK_Q}{TOK_POINT} {tag}{sanitize(spec.instructions).strip()}\n{TOK_ANSWER}\n" | |
| head = {"choice": TOK_CHOICE, "score": TOK_SCORE, "noul": TOK_NOUL}[spec.qtype] | |
| lines = [f"{TOK_Q}{head} {sanitize(spec.instructions).strip()}"] | |
| for i, (name, desc) in enumerate(zip(spec.option_names, spec.option_descs)): | |
| name = sanitize(str(name)).strip() | |
| if desc: | |
| lines.append(f"{TOK_OPT[i]} {name}: {sanitize(str(desc)).strip()}") | |
| else: | |
| lines.append(f"{TOK_OPT[i]} {name}") | |
| lines.append(TOK_ANSWER) | |
| return "\n".join(lines) + "\n" | |
| class Encoded: | |
| input_ids: List[int] | |
| answer_positions: List[int] # index of each <|ts_answer|> token, question order | |
| option_counts: List[int] # k per question (valid slots 0..k-1) | |
| specs: List[QuestionSpec] | |
| labels: List[int] = field(default_factory=list) # -100 if unknown | |
| truncated_state: bool = False | |
| # vision variant | |
| image_ids: List[str] = field(default_factory=list) | |
| image_grid_thw: List[List[int]] = field(default_factory=list) # per image (t, h, w) in patches | |
| image_spans: List[Tuple[int, int]] = field(default_factory=list) # [start, end) positions of each image's tokens | |
| image_sizes: List[Tuple[int, int]] = field(default_factory=list) # original (W, H) | |
| pixel_values: Any = None # torch.Tensor (sum patches, C*T*P*P) or None | |
| merge_size: int = 2 | |
| point_image_index: List[int] = field(default_factory=list) # per question: image index or -1 | |
| point_targets: List[Any] = field(default_factory=list) # per question: tensor (n_tokens+1) or None | |
| def n_tokens(self) -> int: | |
| return len(self.input_ids) | |
| def n_visual_tokens(self) -> int: | |
| return sum(e - s for s, e in self.image_spans) | |
| def merged_grid(self, i: int) -> Tuple[int, int]: | |
| t, h, w = self.image_grid_thw[i] | |
| return h // self.merge_size, w // self.merge_size | |
| def n_tokens_of_image(self, i: int) -> int: | |
| s, e = self.image_spans[i] | |
| return e - s | |
| class Formatter: | |
| """Tokenizer-aware encoder shared by training and inference.""" | |
| def __init__( | |
| self, | |
| tokenizer, | |
| *, | |
| max_total_tokens: int = DEFAULT_MAX_TOTAL_TOKENS, | |
| max_state_tokens: int = DEFAULT_MAX_STATE_TOKENS, | |
| image_processor=None, | |
| max_pixels: Optional[int] = None, | |
| ): | |
| self.tok = tokenizer | |
| self.image_processor = image_processor | |
| self.vision = image_processor is not None | |
| add_special_tokens(self.tok, vision=self.vision) | |
| self.max_total_tokens = max_total_tokens | |
| self.max_state_tokens = max_state_tokens | |
| self.max_pixels = max_pixels | |
| self.image_device = None # set by the client: patchify maths runs on the model device | |
| self.answer_id = self.tok.convert_tokens_to_ids(TOK_ANSWER) | |
| self.state_id = self.tok.convert_tokens_to_ids(TOK_STATE) | |
| self.opt_ids = self.tok.convert_tokens_to_ids(TOK_OPT) | |
| assert self.answer_id is not None and self.answer_id != self.tok.unk_token_id | |
| if self.vision: | |
| self.image_pad_id = self.tok.convert_tokens_to_ids(QWEN_IMAGE_PAD) | |
| self.vision_start_id = self.tok.convert_tokens_to_ids(QWEN_VISION_START) | |
| self.vision_end_id = self.tok.convert_tokens_to_ids(QWEN_VISION_END) | |
| self.point_id = self.tok.convert_tokens_to_ids(TOK_POINT) | |
| for v in (self.image_pad_id, self.vision_start_id, self.vision_end_id, self.point_id): | |
| assert v is not None and v != self.tok.unk_token_id, "tokenizer lacks the Qwen vision tokens" | |
| def _image_block(self, ref, n_tokens: int, size: Tuple[int, int]) -> List[int]: | |
| """<|ts_img|> [image:id WxH role]\n<|vision_start|> pad*n <|vision_end|>\n -> ids; the pad run is contiguous.""" | |
| role = f" {sanitize(str(ref.role))}" if getattr(ref, "role", None) else "" | |
| head = self._ids(f"{TOK_IMG} [image:{sanitize(ref.id)} {size[0]}x{size[1]}{role}]\n") | |
| return head + [self.vision_start_id] + [self.image_pad_id] * n_tokens + [self.vision_end_id] + self._ids("\n") | |
| def _ids(self, text: str) -> List[int]: | |
| return self.tok(text, add_special_tokens=False)["input_ids"] | |
| def encode( | |
| self, | |
| state: Union[str, Dict[str, Any], List[Any]], | |
| questions: Dict[str, Question], | |
| *, | |
| labels: Optional[Dict[str, Union[str, int, bool]]] = None, | |
| shuffle_options: bool = False, | |
| shuffle_questions: bool = False, | |
| drop_label_for: Optional[Sequence[str]] = None, | |
| rng: Optional[random.Random] = None, | |
| state_indent: Optional[int] = None, | |
| option_orders: Optional[Dict[str, Sequence[int]]] = None, | |
| images: Optional[Sequence[Any]] = None, | |
| processed_images=None, | |
| ) -> Encoded: | |
| """images: list of ImageRef / dicts (vision variant). processed_images: a ProcessedImages to reuse (e.g. across | |
| the option-order permutations of one request) instead of running the image processor again.""" | |
| rng = rng or random.Random() | |
| labels = labels or {} | |
| drop = set(drop_label_for or []) | |
| option_orders = option_orders or {} | |
| qids = list(questions.keys()) | |
| if shuffle_questions: | |
| rng.shuffle(qids) | |
| # ---- images (vision variant) | |
| proc = processed_images | |
| refs: List[Any] = [] | |
| if images: | |
| if not self.vision: | |
| raise ValueError("this model is text-only; images and point questions need OpenThai-SystemOne-Vision") | |
| from .images import ImageRef, process_images | |
| refs = [ImageRef.parse(im) for im in images] | |
| if proc is None: | |
| kw = {"max_pixels": self.max_pixels} if self.max_pixels else {} | |
| proc = process_images(self.image_processor, refs, device=self.image_device, **kw) | |
| has_point = any(isinstance(questions[q], Point) for q in qids) | |
| if has_point and not refs: | |
| raise ValueError("point questions need at least one image") | |
| image_index = {r.id: i for i, r in enumerate(refs)} | |
| specs = [ | |
| question_to_spec( | |
| qid, | |
| questions[qid], | |
| label=labels.get(qid), | |
| shuffle=shuffle_options, | |
| rng=rng, | |
| drop_label=qid in drop, | |
| perm=option_orders.get(qid), | |
| ) | |
| for qid in qids | |
| ] | |
| point_image_index: List[int] = [] | |
| for sp in specs: | |
| if sp.qtype != "point": | |
| point_image_index.append(-1) | |
| continue | |
| if sp.image is None: | |
| point_image_index.append(0) | |
| elif sp.image in image_index: | |
| point_image_index.append(image_index[sp.image]) | |
| else: | |
| raise ValueError(f"point question {sp.qid!r} references unknown image {sp.image!r}") | |
| q_texts = [spec_to_text(s, image_id=(refs[point_image_index[i]].id if s.qtype == "point" else None)) for i, s in enumerate(specs)] | |
| q_ids = [self._ids(t) for t in q_texts] | |
| q_total = sum(len(x) for x in q_ids) | |
| # image blocks: inline where the state text says <image:id>, otherwise all before the state | |
| blocks: Dict[str, List[int]] = {} | |
| if refs: | |
| for i, r in enumerate(refs): | |
| blocks[r.id] = self._image_block(r, proc.n_tokens[i], proc.sizes[i]) | |
| state_text = sanitize(state_to_text(state, indent=state_indent)).strip() | |
| inline = [r.id for r in refs if f"<image:{r.id}>" in state_text] | |
| prefix_ids: List[int] = [] | |
| for r in refs: | |
| if r.id not in inline: | |
| prefix_ids += blocks[r.id] | |
| body_parts: List[List[int]] = [] | |
| if inline: | |
| import re | |
| pieces = re.split("(" + "|".join(re.escape(f"<image:{i}>") for i in inline) + ")", state_text) | |
| first = True | |
| for piece in pieces: | |
| if piece.startswith("<image:") and piece[7:-1] in blocks: | |
| body_parts.append(self._ids("\n") + blocks[piece[7:-1]]) | |
| elif piece: | |
| body_parts.append(self._ids((TOK_STATE + " " if first else "") + piece)) | |
| first = False | |
| if first: | |
| body_parts.insert(0, self._ids(TOK_STATE + " ")) | |
| state_ids = [t for part in body_parts for t in part] + self._ids("\n") | |
| else: | |
| state_ids = self._ids(TOK_STATE + " " + state_text + "\n") | |
| budget = min(self.max_state_tokens, self.max_total_tokens - q_total - len(prefix_ids)) | |
| truncated = False | |
| if len(state_ids) > budget and not inline: | |
| # keep the head (state token) and the tail of the state; the end is usually the most recent info | |
| keep_tail = max(budget - 1, 0) | |
| state_ids = state_ids[:1] + state_ids[len(state_ids) - keep_tail :] | |
| truncated = True | |
| ids: List[int] = prefix_ids + list(state_ids) | |
| # locate each image's contiguous pad run (in `refs` order = the order the blocks were emitted) | |
| image_spans: List[Tuple[int, int]] = [] | |
| if refs: | |
| runs = [] | |
| i = 0 | |
| while i < len(ids): | |
| if ids[i] == self.image_pad_id: | |
| j = i | |
| while j < len(ids) and ids[j] == self.image_pad_id: | |
| j += 1 | |
| runs.append((i, j)) | |
| i = j | |
| else: | |
| i += 1 | |
| # runs appear in emission order: prefix blocks first (refs order minus inline), then inline ones in text order | |
| order = [r.id for r in refs if r.id not in inline] + [m for m in re.findall(r"<image:([^>]+)>", state_text) if m in inline] if inline else [r.id for r in refs] | |
| by_id = dict(zip(order, runs)) | |
| image_spans = [by_id[r.id] for r in refs] | |
| assert all(e - s == proc.n_tokens[i] for i, (s, e) in enumerate(image_spans)) | |
| answer_positions: List[int] = [] | |
| for qi in q_ids: | |
| ids.extend(qi) | |
| # the answer token is the last non-newline token of each question block | |
| pos = len(ids) - 1 | |
| while ids[pos] != self.answer_id: | |
| pos -= 1 | |
| answer_positions.append(pos) | |
| point_targets: List[Any] = [] | |
| if refs: | |
| from .images import bbox_to_token_target | |
| for sp, ii in zip(specs, point_image_index): | |
| if sp.qtype == "point" and sp.point_label_given: | |
| t, h, w = proc.grid_thw[ii] | |
| point_targets.append(bbox_to_token_target(sp.point_bbox, h // proc.merge, w // proc.merge)) | |
| else: | |
| point_targets.append(None) | |
| else: | |
| point_targets = [None] * len(specs) | |
| return Encoded( | |
| input_ids=ids, | |
| answer_positions=answer_positions, | |
| option_counts=[len(s.option_names) for s in specs], | |
| specs=specs, | |
| labels=[(-100 if s.label_slot is None else s.label_slot) for s in specs], | |
| truncated_state=truncated, | |
| image_ids=[r.id for r in refs], | |
| image_grid_thw=list(proc.grid_thw) if refs else [], | |
| image_spans=image_spans, | |
| image_sizes=list(proc.sizes) if refs else [], | |
| pixel_values=proc.pixel_values if refs else None, | |
| merge_size=proc.merge if refs else 2, | |
| point_image_index=point_image_index, | |
| point_targets=point_targets, | |
| ) | |
| def slot_mask(option_counts: Sequence[int], *, include_abstain: bool = True, n_slots: int = N_SLOTS): | |
| """Boolean mask (Q, n_slots): True where a slot is valid for that question.""" | |
| import torch | |
| k = torch.as_tensor(list(option_counts), dtype=torch.long) | |
| ar = torch.arange(n_slots) | |
| mask = ar[None, :] < k[:, None] | |
| if include_abstain: | |
| mask[:, ABSTAIN_SLOT] = True | |
| return mask | |
| def mrope_position_ids(encoded: Sequence[Encoded], T: int): | |
| """Qwen3.5 3-D (t, h, w) position ids for a right-padded batch, computed on CPU without per-token python loops. | |
| Text tokens: t = h = w = running position. An image with merged grid (t, hm, wm) placed at running position p gets | |
| t = p + frame index, h = p + row, w = p + col, and advances the running position by max(hm, wm). Equals | |
| `Qwen3_5Model.get_rope_index` (tests/test_vision.py checks it) but costs ~0.1 ms instead of tens of ms. | |
| """ | |
| import torch | |
| B = len(encoded) | |
| pos = torch.zeros((3, B, T), dtype=torch.long) | |
| for b, e in enumerate(encoded): | |
| n = e.n_tokens | |
| cur = 0 # running position | |
| idx = 0 # token index | |
| spans = sorted(zip(e.image_spans, range(len(e.image_spans)))) | |
| for (s0, e0), i in spans: | |
| if s0 > idx: # text before the image | |
| ar = torch.arange(s0 - idx) + cur | |
| pos[:, b, idx:s0] = ar | |
| cur += s0 - idx | |
| t, h, w = e.image_grid_thw[i] | |
| hm, wm = h // e.merge_size, w // e.merge_size | |
| tt = torch.arange(t).view(t, 1, 1).expand(t, hm, wm) | |
| hh = torch.arange(hm).view(1, hm, 1).expand(t, hm, wm) | |
| ww = torch.arange(wm).view(1, 1, wm).expand(t, hm, wm) | |
| pos[0, b, s0:e0] = tt.reshape(-1) + cur | |
| pos[1, b, s0:e0] = hh.reshape(-1) + cur | |
| pos[2, b, s0:e0] = ww.reshape(-1) + cur | |
| cur += max(hm, wm) | |
| idx = e0 | |
| if n > idx: | |
| pos[:, b, idx:n] = torch.arange(n - idx) + cur | |
| return pos | |
| def collate(encoded: Sequence[Encoded], pad_id: int, *, max_questions: Optional[int] = None): | |
| """Right-pad a batch. Returns dict of tensors for OpenThaiSystemOneForDecision.forward.""" | |
| import torch | |
| B = len(encoded) | |
| T = max(e.n_tokens for e in encoded) | |
| Q = max_questions or max(len(e.answer_positions) for e in encoded) | |
| input_ids = torch.full((B, T), pad_id, dtype=torch.long) | |
| attention_mask = torch.zeros((B, T), dtype=torch.long) | |
| answer_positions = torch.zeros((B, Q), dtype=torch.long) | |
| option_counts = torch.zeros((B, Q), dtype=torch.long) | |
| labels = torch.full((B, Q), -100, dtype=torch.long) | |
| qtypes = torch.full((B, Q), -1, dtype=torch.long) | |
| for b, e in enumerate(encoded): | |
| n = e.n_tokens | |
| input_ids[b, :n] = torch.tensor(e.input_ids) | |
| attention_mask[b, :n] = 1 | |
| q = len(e.answer_positions) | |
| answer_positions[b, :q] = torch.tensor(e.answer_positions) | |
| option_counts[b, :q] = torch.tensor(e.option_counts) | |
| qtypes[b, :q] = torch.tensor([QTYPES.index(s.qtype) for s in e.specs]) | |
| if e.labels: | |
| labels[b, :q] = torch.tensor(e.labels) | |
| out = { | |
| "input_ids": input_ids, | |
| "attention_mask": attention_mask, | |
| "answer_positions": answer_positions, | |
| "option_counts": option_counts, | |
| "labels": labels, | |
| "qtypes": qtypes, | |
| } | |
| if any(e.pixel_values is not None for e in encoded): | |
| I = max(len(e.image_spans) for e in encoded) | |
| pv = [e.pixel_values for e in encoded if e.pixel_values is not None] | |
| grid = [torch.tensor(e.image_grid_thw, dtype=torch.long) for e in encoded if e.image_grid_thw] | |
| image_spans = torch.full((B, I, 2), -1, dtype=torch.long) | |
| mm_token_type_ids = torch.zeros((B, T), dtype=torch.long) | |
| point_image_index = torch.full((B, Q), -1, dtype=torch.long) | |
| L = max([e.n_tokens_of_image(i) for e in encoded for i in range(len(e.image_spans))] + [1]) | |
| point_targets = torch.zeros((B, Q, L + 1)) | |
| point_has_target = torch.zeros((B, Q), dtype=torch.bool) | |
| for b, e in enumerate(encoded): | |
| for i, (s0, e0) in enumerate(e.image_spans): | |
| image_spans[b, i] = torch.tensor([s0, e0]) | |
| mm_token_type_ids[b, s0:e0] = 1 | |
| for qi, ii in enumerate(e.point_image_index): | |
| point_image_index[b, qi] = ii | |
| tgt = e.point_targets[qi] if qi < len(e.point_targets) else None | |
| if tgt is not None: | |
| n = tgt.shape[0] - 1 | |
| point_targets[b, qi, :n] = tgt[:n] | |
| point_targets[b, qi, L] = tgt[n] | |
| point_has_target[b, qi] = True | |
| out.update({ | |
| "pixel_values": torch.cat(pv, 0), | |
| "image_grid_thw": torch.cat(grid, 0), | |
| "mm_token_type_ids": mm_token_type_ids, | |
| "position_ids": mrope_position_ids(encoded, T), | |
| "image_spans": image_spans, | |
| "point_image_index": point_image_index, | |
| "point_targets": point_targets, | |
| "point_has_target": point_has_target, | |
| }) | |
| return out | |