dmis-lab commited on
Commit
a381a62
·
verified ·
1 Parent(s): a8dd5fa

Add files using upload-large-folder tool

Browse files
Files changed (40) hide show
  1. NOTICE.txt +4 -0
  2. README.md +121 -0
  3. config.json +30 -0
  4. configuration_cvrr_merged.py +18 -0
  5. configuration_source_qwen25.py +1142 -0
  6. configuration_source_qwen3.py +740 -0
  7. cvrr_release_config.json +19 -0
  8. licenses/Apache-2.0.txt +202 -0
  9. licenses/InternVL-MIT.txt +21 -0
  10. licenses/UPSTREAM_MODEL_CARD.md +701 -0
  11. merged_transition.safetensors +3 -0
  12. modeling_cvrr_merged.py +279 -0
  13. modeling_source_qwen25.py +0 -0
  14. modeling_source_qwen3.py +0 -0
  15. native_backbone/added_tokens.json +11 -0
  16. native_backbone/config.json +225 -0
  17. native_backbone/configuration_intern_vit.py +119 -0
  18. native_backbone/configuration_internlm2.py +150 -0
  19. native_backbone/configuration_internvl_chat.py +101 -0
  20. native_backbone/conversation.py +391 -0
  21. native_backbone/generation_config.json +4 -0
  22. native_backbone/model.safetensors.index.json +692 -0
  23. native_backbone/modeling_intern_vit.py +429 -0
  24. native_backbone/modeling_internlm2.py +1456 -0
  25. native_backbone/modeling_internvl_chat.py +363 -0
  26. native_backbone/native-00001.safetensors +3 -0
  27. native_backbone/native-00002.safetensors +3 -0
  28. native_backbone/native-00003.safetensors +3 -0
  29. native_backbone/native-00004.safetensors +3 -0
  30. native_backbone/special_tokens_map.json +63 -0
  31. native_backbone/tokenization_internlm3.py +294 -0
  32. native_backbone/tokenizer.model +3 -0
  33. native_backbone/tokenizer_config.json +330 -0
  34. requirements.txt +11 -0
  35. source_gemma.py +685 -0
  36. source_helpers.py +69 -0
  37. source_internvl.py +286 -0
  38. source_perceive.py +414 -0
  39. source_spatial.py +349 -0
  40. source_splitting.py +473 -0
NOTICE.txt ADDED
@@ -0,0 +1,4 @@
 
 
 
 
 
1
+ CVRR applies a learned recurrent transition to OpenGVLab/InternVL3-9B.
2
+ Native dense weights are preserved; recurrent projection weights are modified
3
+ by merging the trained CVRR LoRA. CVRR-specific code is additional code, not
4
+ part of the original backbone distribution. See the included upstream notices.
README.md ADDED
@@ -0,0 +1,121 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ base_model: OpenGVLab/InternVL3-9B
3
+ base_model_relation: finetune
4
+ library_name: transformers
5
+ license: mit
6
+ tags:
7
+ - cvrr
8
+ - custom_code
9
+ - latent-reasoning
10
+ - visual-question-answering
11
+ - arxiv:2609.06746
12
+ ---
13
+
14
+ # InternVL3-9B-CVRR
15
+
16
+ <p align="center">
17
+ <a href="https://arxiv.org/abs/2609.06746"><strong>📄 Paper</strong></a> ·
18
+ <a href="https://github.com/dmis-lab/CVRR"><strong>💻 Code</strong></a>
19
+ </p>
20
+
21
+ ## Introduction
22
+
23
+ Latent visual reasoning is not just about storing visual information in hidden
24
+ states. That information should matter for the model's answer. When an answer
25
+ decoder can access the original multimodal context through a parallel path,
26
+ the presence of informative latent states alone does not establish their use.
27
+
28
+ Our work, **Reason Through the Latent! Making Latent Visual Reasoning Necessary**,
29
+ introduces **Causal Visual Recurrent Reasoning (CVRR)**. CVRR separates visual
30
+ access during reasoning from visual access during answering, focusing on three
31
+ design principles:
32
+
33
+ **👁️ Preserving pretrained visual competence**
34
+ Reasoning starts from the native image-conditioned question representation,
35
+ preserving visual information already integrated by the pretrained VLM.
36
+
37
+ **🔁 Refining reasoning with persistent visual evidence**
38
+ A shared native decoder layer updates the question state while retaining access
39
+ to fixed visual evidence. The recurrent transition is adapted with LoRA and
40
+ answer-token supervision.
41
+
42
+ **🔒 Making the reasoning state the visual answer interface**
43
+ The upper answer decoder receives the final recurrent question state, without
44
+ direct access to the original visual rows or original multimodal prefix caches.
45
+ The lower-layer prefix context used for answer generation follows a text-only path.
46
+
47
+ ## Model Checkpoints
48
+
49
+ ### Final Models
50
+
51
+ We release CVRR checkpoints based on the following pretrained vision-language
52
+ backbones. Each repository contains the native backbone weights, the merged
53
+ recurrent transition, and the custom inference code; no separate LoRA download
54
+ is required.
55
+
56
+ - 🤗 [Qwen2.5-VL-7B-CVRR](https://huggingface.co/dmis-lab/Qwen2.5-VL-7B-CVRR)
57
+ - 🤗 [Qwen3-VL-8B-CVRR](https://huggingface.co/dmis-lab/Qwen3-VL-8B-CVRR)
58
+ - 🤗 [Gemma3-4B-CVRR](https://huggingface.co/dmis-lab/Gemma3-4B-CVRR)
59
+ - 🤗 [Gemma3-12B-CVRR](https://huggingface.co/dmis-lab/Gemma3-12B-CVRR)
60
+ - 🤗 [Gemma4-12B-CVRR](https://huggingface.co/dmis-lab/Gemma4-12B-CVRR)
61
+ - 🤗 [InternVL3-9B-CVRR](https://huggingface.co/dmis-lab/InternVL3-9B-CVRR)
62
+
63
+ This checkpoint is based on **`OpenGVLab/InternVL3-9B`**.
64
+
65
+ ## Loading
66
+
67
+ Download the complete repository and install its model-specific dependencies:
68
+
69
+ ```bash
70
+ hf download dmis-lab/InternVL3-9B-CVRR --local-dir ./InternVL3-9B-CVRR
71
+ pip install -r ./InternVL3-9B-CVRR/requirements.txt
72
+ ```
73
+
74
+ ```python
75
+ import torch
76
+ from PIL import Image
77
+ from transformers import AutoModelForImageTextToText
78
+
79
+ model = AutoModelForImageTextToText.from_pretrained(
80
+ "./InternVL3-9B-CVRR",
81
+ trust_remote_code=True,
82
+ dtype=torch.bfloat16,
83
+ device_map="cuda:0",
84
+ ).eval()
85
+
86
+ image = Image.open("example.jpg").convert("RGB")
87
+ inputs = model.prepare_inputs(
88
+ image,
89
+ "What is the dominant color? A. Red B. Blue C. Green D. Yellow. Answer with the option letter.",
90
+ )
91
+ logits = model.next_token_logits(**inputs)
92
+ token_id = logits.argmax(dim=-1).item()
93
+ print(model.tokenizer.decode([token_id]))
94
+ ```
95
+
96
+ This model exposes the strict first-answer-token readout used in our analyses. Multi-token generation is not currently supported by this loader.
97
+
98
+ Use one complete model replica per GPU. This custom Transformers loader does
99
+ not support automatic model sharding, generic `save_pretrained()` reserialization,
100
+ or direct loading through vLLM/SGLang. Preserve the downloaded directory and
101
+ pin the Hub revision for reproducible use.
102
+
103
+ ## License
104
+
105
+ Please follow the upstream model's terms and the included
106
+ [`NOTICE.txt`](NOTICE.txt) and [`licenses/`](licenses/) attribution files.
107
+ The component licenses remain applicable to their respective materials.
108
+
109
+ ## Citation
110
+
111
+ ```bibtex
112
+ @misc{park2026reasonlatentmakinglatent,
113
+ title={Reason Through the Latent! Making Latent Visual Reasoning Necessary},
114
+ author={Suhyeong Park and Junha Jung and Jaewoo Kang},
115
+ year={2026},
116
+ eprint={2609.06746},
117
+ archivePrefix={arXiv},
118
+ primaryClass={cs.AI},
119
+ url={https://arxiv.org/abs/2609.06746},
120
+ }
121
+ ```
config.json ADDED
@@ -0,0 +1,30 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "model_type": "cvrr_merged",
3
+ "architectures": [
4
+ "CVRRMergedModel"
5
+ ],
6
+ "auto_map": {
7
+ "AutoConfig": "configuration_cvrr_merged.CVRRMergedConfig",
8
+ "AutoModel": "modeling_cvrr_merged.CVRRMergedModel",
9
+ "AutoModelForImageTextToText": "modeling_cvrr_merged.CVRRMergedModel"
10
+ },
11
+ "release": {
12
+ "name": "CVRR-InternVL3-9B",
13
+ "base_model": "OpenGVLab/InternVL3-9B",
14
+ "ell_star": 35,
15
+ "recurrent_layer": 36,
16
+ "upper_decoder_start": 37,
17
+ "inference_T": 4,
18
+ "inference_beta": 0.33,
19
+ "lora_rank": 32,
20
+ "lora_alpha": 12.0,
21
+ "lora_dropout": 0.01,
22
+ "checkpoint_step": 500,
23
+ "format": "cvrr_native_plus_merged_transition_v1",
24
+ "merged_weight_dtype": "float32",
25
+ "native_weight_bytes": 18277586944,
26
+ "status": "gpu_vstar_comparison_completed",
27
+ "upload_ready": true,
28
+ "hub_model_id": "dmis-lab/InternVL3-9B-CVRR"
29
+ }
30
+ }
configuration_cvrr_merged.py ADDED
@@ -0,0 +1,18 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from transformers import PretrainedConfig
2
+
3
+
4
+ class CVRRMergedConfig(PretrainedConfig):
5
+ model_type = 'cvrr_merged'
6
+
7
+ def __init__(self, release=None, **kwargs):
8
+ super().__init__(**kwargs)
9
+ self.release = {} if release is None else release
10
+ self.architectures = ['CVRRMergedModel']
11
+ self.auto_map = {
12
+ 'AutoConfig': 'configuration_cvrr_merged.CVRRMergedConfig',
13
+ 'AutoModel': 'modeling_cvrr_merged.CVRRMergedModel',
14
+ 'AutoModelForImageTextToText': 'modeling_cvrr_merged.CVRRMergedModel',
15
+ }
16
+
17
+
18
+ CVRRMergedConfig.register_for_auto_class()
configuration_source_qwen25.py ADDED
@@ -0,0 +1,1142 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Configuration for CLOSE latent reasoning on Qwen2.5-VL.
2
+
3
+ The current ``raw_evidence_question`` path keeps the complete layer-``ell_star``
4
+ visual memory as immutable evidence and recurrently updates the text-only
5
+ question state by cross-attending to it. Its selected interface averages all
6
+ four residual recurrent states and may use a training-only restoration target
7
+ from the normal multimodal question state. Historical slot-workspace fields are
8
+ retained so old checkpoints remain loadable, but they are mutually exclusive
9
+ with the raw-evidence path.
10
+
11
+ Every architectural hyperparameter of Method §3 lives here -- nothing is
12
+ hardcoded in ``modeling_close_qwen2_5_vl.py``.
13
+
14
+ Layer-index convention (used consistently across the repo)
15
+ ----------------------------------------------------------
16
+ ``ell_star`` is **0-indexed and inclusive**: it is the index of the last decoder
17
+ layer belonging to the lower branch. With ``num_hidden_layers == 28``::
18
+
19
+ F_{<=l*} = language_model.layers[: ell_star + 1] # lower_slice
20
+ F_{>l*} = language_model.layers[ell_star + 1 :] # upper_slice
21
+
22
+ so ``ell_star=18`` means layers 0..18 run before the split and 19..27 after it.
23
+ ``ell_star`` is measured once by ``scripts/localize_read_point.py`` (§3.1) and
24
+ frozen before any workspace training.
25
+
26
+ Careful with ``Qwen2_5_VLConfig``: its ``__setattr__`` forwards any attribute
27
+ whose name already exists in ``text_config.__dict__`` down to the sub-config, and
28
+ its ``__init__`` builds ``text_config`` from ``**kwargs``. Workspace fields are
29
+ therefore declared as explicit named parameters so they never enter ``kwargs``
30
+ and never shadow a text-config field. ``tests/test_config_roundtrip.py`` guards
31
+ this.
32
+ """
33
+
34
+ from __future__ import annotations
35
+
36
+ import math
37
+
38
+ from transformers.models.qwen2_5_vl.configuration_qwen2_5_vl import Qwen2_5_VLConfig
39
+
40
+ #: How the ``S`` workspace slots are assigned M-RoPE positions once they are
41
+ #: spliced into the replacement cache above ``ell_star``. Qwen2.5-VL uses 3D
42
+ #: M-RoPE (t, h, w) with ``mrope_section=[16, 24, 24]``, so slots -- which have
43
+ #: no spatial extent -- need an explicit convention.
44
+ WORKSPACE_ROPE_MODES = (
45
+ # Slots continue the 1-D text position sequence directly after Q*, with
46
+ # t == h == w. Treats the workspace as "more text".
47
+ "continue",
48
+ # Slots reuse the (t, h, w) span the image tokens occupied in the original
49
+ # multimodal prefill, subsampled to S positions. Preserves the read point's
50
+ # positional geometry but reintroduces image-derived indices.
51
+ "image_span",
52
+ # All slots share one position (the first position after Q*). Makes the
53
+ # workspace order-free, matching the set semantics of Perceiver-style slots.
54
+ "shared",
55
+ )
56
+
57
+
58
+ class CloseQwen2_5_VLConfig(Qwen2_5_VLConfig):
59
+ r"""Configuration for :class:`CloseQwen2_5_VLForConditionalGeneration`.
60
+
61
+ Args:
62
+ ell_star: 0-indexed inclusive last layer of the lower branch. ``None``
63
+ until §3.1 localization has run; the model refuses to build with
64
+ ``None`` rather than falling back to a guessed depth.
65
+ num_workspace_steps: the fixed recurrent horizon (paper notation ``T``;
66
+ per-item segment counts ``n`` exist only in the losses, never here).
67
+ The trainable parameter count is independent of it because
68
+ ``f_theta`` is shared across steps, so it is a free experimental
69
+ axis. ``0`` is legal and is the no-recurrence floor: ``Z^(0)`` goes
70
+ straight to the decoder.
71
+ num_workspace_slots: ``S``, slots in ``Z``. Free capacity knob.
72
+ workspace_width: ``d_w``, the slot width. Free capacity knob -- the
73
+ method's claim is path exclusivity (all visual information routed
74
+ through ``Z``), NOT an information bottleneck, so ``S`` and ``d_w``
75
+ carry no methodological commitment. Kept below ``d`` (3584 for the
76
+ 7B) since ``P_Z`` projects up at the splice.
77
+ workspace_num_heads: attention heads inside ``r_theta`` / ``f_theta``.
78
+ read_num_blocks: cross-attention blocks in ``r_theta``. The *read* is
79
+ once regardless of depth: all blocks see the same ``V*`` in a single
80
+ forward, and ``V*`` is discarded afterwards.
81
+ transition_num_blocks: blocks in the shared recurrent ``f_theta``.
82
+ workspace_ffn_mult: FFN expansion inside workspace blocks.
83
+ workspace_dropout: dropout inside workspace blocks.
84
+ workspace_rope_mode: one of :data:`WORKSPACE_ROPE_MODES`.
85
+ adapter_rank: rank of the low-rank adapter on the post-``ell_star``
86
+ backbone layers. ``0`` disables it (frozen backbone only).
87
+ adapter_alpha: LoRA-style scaling; effective scale is
88
+ ``adapter_alpha / adapter_rank``.
89
+ adapter_dropout: dropout probability applied only to the LoRA branch.
90
+ lambda_traj: legacy trajectory-loss weight. The current method sets it
91
+ to zero and trains with final-answer CE only; it remains serialized
92
+ solely so historical checkpoints can still be inspected.
93
+ """
94
+
95
+ model_type = "close_qwen2_5_vl"
96
+
97
+ def __init__(
98
+ self,
99
+ ell_star: int | None = None,
100
+ num_workspace_steps: int = 4,
101
+ num_workspace_slots: int = 32,
102
+ workspace_width: int = 512,
103
+ workspace_num_heads: int = 8,
104
+ read_num_blocks: int = 1,
105
+ transition_num_blocks: int = 2,
106
+ workspace_ffn_mult: int = 4,
107
+ workspace_dropout: float = 0.0,
108
+ workspace_rope_mode: str = "continue",
109
+ adapter_rank: int = 16,
110
+ adapter_alpha: int = 32,
111
+ adapter_dropout: float = 0.0,
112
+ adapter_exclude_layers: tuple = (),
113
+ lambda_traj: float = 0.0,
114
+ read_shortcut: bool = False,
115
+ splice_scale_match: bool = False,
116
+ read_gating: bool = False,
117
+ gate_init_eps: float = 0.02,
118
+ read_residual_pure: bool = False,
119
+ residual_norm_cap: float = 0.0,
120
+ pure_warm_start: bool = False,
121
+ final_verify: bool = False,
122
+ mid_read_step: int = 0,
123
+ persistent_slot_id: bool = False,
124
+ mid_read_alpha: float = 1.0,
125
+ evidence_reasoning: bool = False,
126
+ raw_evidence_question: bool = False,
127
+ raw_attention_source_layer: int = -1,
128
+ raw_transition_ffn: bool | None = None,
129
+ raw_state_replacement: bool = False,
130
+ raw_question_anchor: bool = False,
131
+ raw_mean_aggregation: bool = False,
132
+ hierarchical_visual_layers: tuple = (),
133
+ counterfactual_residual_recurrence: bool = False,
134
+ visual_counterfactual_recurrence: bool = False,
135
+ visual_cumulative_recurrence: bool = False,
136
+ visual_full_state_recurrence: bool = False,
137
+ visual_boundary_reentry: bool = False,
138
+ visual_source_centered_reentry: bool = False,
139
+ recurrence_only_adapter: bool = False,
140
+ visual_preserving_adapter_correction: bool = False,
141
+ perceive_deliberate_chain: bool = False,
142
+ local_visual_latent_chain: bool = False,
143
+ pdlc_transition_conditioned: bool = False,
144
+ pdlc_question_visible_through_pair: int = 1,
145
+ pdlc_question_curriculum: bool = False,
146
+ latent_chain_init_token_id: int = -1,
147
+ spatial_visual_recurrence: bool = False,
148
+ spatial_recurrence_steps: int = 8,
149
+ spatial_recurrence_inner_width: int = 256,
150
+ spatial_recurrence_heads: int = 8,
151
+ spatial_step_override: int = -1,
152
+ counterfactual_beta: float = 1.0,
153
+ crr_concat_aggregation: bool = False,
154
+ crr_policy_rank: int = 0,
155
+ crr_policy_alpha: float = 32.0,
156
+ counterfactual_reliance_loss: bool = False,
157
+ transition_aware_reliance_loss: bool = False,
158
+ corrective_transition_reliance_loss: bool = False,
159
+ counterfactual_step_retention_loss: bool = False,
160
+ functional_state_loss: bool = False,
161
+ per_example_answer_loss: bool = False,
162
+ interface_loss: bool = False,
163
+ read_anchor: bool = False,
164
+ **kwargs,
165
+ ):
166
+ super().__init__(**kwargs)
167
+ self.ell_star = ell_star
168
+ self.num_workspace_steps = num_workspace_steps
169
+ self.num_workspace_slots = num_workspace_slots
170
+ self.workspace_width = workspace_width
171
+ self.workspace_num_heads = workspace_num_heads
172
+ self.read_num_blocks = read_num_blocks
173
+ self.transition_num_blocks = transition_num_blocks
174
+ self.workspace_ffn_mult = workspace_ffn_mult
175
+ self.workspace_dropout = workspace_dropout
176
+ self.workspace_rope_mode = workspace_rope_mode
177
+ self.adapter_rank = adapter_rank
178
+ self.adapter_alpha = adapter_alpha
179
+ self.adapter_dropout = adapter_dropout
180
+ # E run: dense full-FT 층은 LoRA 대상에서 제외 (이중 파라미터화 방지)
181
+ self.adapter_exclude_layers = list(adapter_exclude_layers or ())
182
+ self.lambda_traj = lambda_traj
183
+ # v2 anti-collapse switches (defaults False so v1 checkpoints reload).
184
+ # read_shortcut: r_theta adds an orthogonally-projected pooled-V* term to
185
+ # every slot, so Z^(0) is image-dependent from step 0 instead of relying
186
+ # on a random cross-attention finding signal.
187
+ # splice_scale_match: P_Z(Z) rows are RMS-matched to the question rows at
188
+ # the splice. Unmatched, slot states are ~O(1) against layer-20 hiddens
189
+ # in the hundreds -- the decoder could ignore them for free, and did
190
+ # (probe 20143: visual sensitivity 0.94 -> 1.0000 over training).
191
+ self.read_shortcut = read_shortcut
192
+ self.splice_scale_match = splice_scale_match
193
+ # v6 candidate: gated re-reads DURING the recurrence. Every re-read
194
+ # goes through r_theta into Z, so the mediation invariant (decoder sees
195
+ # only [Q*; P_Z(Z)]) is untouched; what changes is that the workspace
196
+ # may consult V* again mid-reasoning, scaled by a learned scalar gate.
197
+ # The gate is free at EVERY step -- no hand-coded schedule. Degeneration
198
+ # into answer-time retrieval is disciplined by the objective itself:
199
+ # L_traj pins each intermediate state to its teacher segment, so a
200
+ # wait-then-fetch trajectory misses its intermediate targets. The gate
201
+ # bias initialises negative, i.e. training starts from the proven
202
+ # no-reread regime and departs only where gradients demand it.
203
+ # Default False: v4/v5 checkpoints reload bit-identically.
204
+ self.read_gating = read_gating
205
+ # Initial gate open-rate epsilon_gamma: b_gamma = logit(eps). Sigmoid
206
+ # never reaches exactly 0 at finite logits, so this is near-identity
207
+ # (not bitwise) initialisation w.r.t. the read-once path.
208
+ self.gate_init_eps = gate_init_eps
209
+ # v7: always-on PURE visual residual rereads. Every step adds
210
+ # Z = U + W_o XAttn(LN(U) -> queries, V* -> keys/values): no self-attn,
211
+ # no FFN, no output norm, no biases in the branch, W_o zero-init.
212
+ # Motivation (battery, ep1 ckpt 20752): read_once hurt badly while
213
+ # null_step (blank memory at one step) was harmless -- the gated
214
+ # reader's value was extra recurrent DEPTH, not vision. This branch
215
+ # removes every path that can change Z without V* content, so any
216
+ # benefit it shows IS visual by construction.
217
+ self.read_residual_pure = read_residual_pure
218
+ # v7.1 (advisor): relative-norm trust region on the visual residual.
219
+ # s_k = min(1, rho * RMS(U) / (RMS(Delta) + eps)); Z = U + s_k * Delta
220
+ # Unlike a fixed alpha, the optimizer cannot undo this by rescaling
221
+ # W_o (s_k adapts); ||Delta_used|| <~ rho * ||U|| structurally, and
222
+ # V* = 0 => Delta = 0 keeps the zero-memory identity. 0.0 disables.
223
+ self.residual_norm_cap = residual_norm_cap
224
+ # v7.1: initialise pure_attn K/V projections from the initial visual
225
+ # reader's v_in (same [d_w, d] shape) and Q from its first read
226
+ # block's in-proj slice, instead of random xavier -- the ep1 finding
227
+ # was that Q/K/V stayed ~random while only W_o trained.
228
+ self.pure_warm_start = pure_warm_start
229
+ # MAIN METHOD (advisor-confirmed): Final-Verified Recurrent Workspace.
230
+ # Initial read -> PURE latent recurrence (no mid-step visual anything)
231
+ # -> exactly ONE pure visual verification at the end:
232
+ # Z_out = U^(T) + W_o XAttn(LN(U^(T)), V*), alpha = 1, W_o = 0 init.
233
+ # Uses the same bias-free verifier module as read_residual_pure; the
234
+ # two flags differ only in WHERE the branch fires (every step vs once).
235
+ self.final_verify = final_verify
236
+ # MAIN (advisor 최종 확정): Mid-Read Recurrent Workspace. 고정 상수
237
+ # k* = floor(T/2) = 4 (스텝 sweep 산물이 아니라 판독 전후 동수 전이의
238
+ # midpoint)에서 pure visual read 1회:
239
+ # Z^(4) = U^(4) + W_o XAttn(LN(U^(4)), V*), alpha=1, W_o=0 init.
240
+ # 논리적 step 4 == python 루프 index 3 (0-based) -- off-by-one은
241
+ # tests/test_mid_read.py가 상태 분기 위치로 고정한다. 0이면 비활성.
242
+ self.mid_read_step = mid_read_step
243
+ # advisor(2026-08-06): E_slot(=workspace_read.slots)을 매 recurrent
244
+ # step 입력과 mid-read query에 일시 제공. state에는 누적하지 않는다.
245
+ self.persistent_slot_id = persistent_slot_id
246
+ # 고정 read 스케일 Z = U + alpha*W_o A (advisor scale-mismatch 검사).
247
+ # 1.0 = 기존과 비트동일. eval의 gate_overrides alpha는 절대값으로 이를 대체.
248
+ self.mid_read_alpha = mid_read_alpha
249
+ # advisor 최종안: E(persistent evidence) / R(recurrent reasoning) 분리.
250
+ # states = [E^(0), R^(1..T)], 결합 Z^(T)=E+R^(T)는 prefill에서.
251
+ self.evidence_reasoning = evidence_reasoning
252
+ # Current replacement candidate: no learned evidence/reasoning slots.
253
+ # E is the complete raw V* sequence and R0 is the text-only Q* sequence.
254
+ # A single shared CrossAttn transition produces R1..RT. Residual mode
255
+ # decodes RT; the replacement candidate decodes Q*+HT. Cross-attention
256
+ # uses the backbone's native GQA shape and
257
+ # is initialized from the first upper layer by default. Checkpoints
258
+ # created before the CrossAttn-only decision did not serialize
259
+ # raw_transition_ffn; None therefore preserves their historical FFN
260
+ # for faithful evaluation, while every new training launch passes False.
261
+ self.raw_evidence_question = bool(raw_evidence_question)
262
+ self.raw_attention_source_layer = int(raw_attention_source_layer)
263
+ self.raw_transition_ffn = (
264
+ True if raw_transition_ffn is None else bool(raw_transition_ffn)
265
+ )
266
+ # Experimental simplification of the same raw-E/Q path. The residual
267
+ # carrier is replaced, not augmented:
268
+ # H0=Q*, Hk=CrossAttn(H{k-1}, E), decoder_latent=Q*+HT.
269
+ # False preserves all residual raw-E/Q checkpoints and active runs.
270
+ self.raw_state_replacement = bool(raw_state_replacement)
271
+ # Question-anchored replacement keeps the fixed text scaffold in every
272
+ # recurrent query without restoring the old recurrent identity path:
273
+ # Z0=Q*, Zk=Q*+CrossAttn(Z{k-1}, E), decoder_latent=ZT.
274
+ # This replaces the one-time Q*+HT decoder interface above.
275
+ self.raw_question_anchor = bool(raw_question_anchor)
276
+ # Q-shaped recurrent candidate selected by the V* interface oracle:
277
+ # R0=Q*, Rk=R{k-1}+CrossAttn(R{k-1}, V*)
278
+ # R_agg=mean(R1..RT)
279
+ # This is parameter-free and replaces both question-anchored state
280
+ # replacement and final-state-only decoding.
281
+ self.raw_mean_aggregation = bool(raw_mean_aggregation)
282
+ # A non-empty sequence maps the first T-1 recurrent steps to frozen
283
+ # visual-token states from the listed language layers. The final step
284
+ # uses Q* as memory to integrate the staged reads without another image
285
+ # access. One CrossAttention module is shared by all T steps.
286
+ self.hierarchical_visual_layers = [
287
+ int(layer) for layer in (hierarchical_visual_layers or ())
288
+ ]
289
+ # Counterfactual Residual Recurrence (CRR): the image-conditioned and
290
+ # text-only question states are propagated by the same native layer
291
+ # F=L(ell_star+1). Only their counterfactual residual is recurrently
292
+ # refined; there are no slots, V* memory, custom attention modules, or
293
+ # inference-time auxiliary modules. beta is the fixed relaxed-update
294
+ # coefficient; an optional batch-swap reliance loss is training-only.
295
+ self.counterfactual_residual_recurrence = bool(
296
+ counterfactual_residual_recurrence
297
+ )
298
+ # Causal Visual Residual Recurrence (CVRR), implemented as the
299
+ # vision-native CRR transition. At every recurrent step the same
300
+ # native layer is evaluated with and without only Q->vision attention;
301
+ # their difference is the image-caused update. It deliberately
302
+ # reuses CRR's strict Q-shaped decoder/cache interface and mean
303
+ # aggregation, hence the parent CRR flag remains explicit.
304
+ self.visual_counterfactual_recurrence = bool(
305
+ visual_counterfactual_recurrence
306
+ )
307
+ # Minimal cumulative CVRR variant. Unlike the historical anchored
308
+ # mean path, visual updates are accumulated into the running state and
309
+ # the final state is decoded directly:
310
+ # R_k = R_{k-1} + beta * Delta_k^V, decode R_T.
311
+ # One explicit flag controls both choices so a checkpoint cannot
312
+ # silently combine cumulative dynamics with the old mean interface.
313
+ self.visual_cumulative_recurrence = bool(
314
+ visual_cumulative_recurrence
315
+ )
316
+ # Full-state CVRR keeps the native recurrent computation instead of
317
+ # treating the visual on/off intervention difference as the next
318
+ # hidden state:
319
+ # R_tilde_k = Pi_Q F([E; R_{k-1}])
320
+ # R_k = R_{k-1} + beta * (R_tilde_k - R_{k-1}).
321
+ # The visual-off branch remains available to evaluation code, but is
322
+ # not part of this model's train/inference transition.
323
+ self.visual_full_state_recurrence = bool(
324
+ visual_full_state_recurrence
325
+ )
326
+ # Input-grounded closure control. Later recurrent calls rebuild the
327
+ # native post-L20 scaffold and inject only the carried question-state
328
+ # deviation, rather than feeding post-L21 rows back into L21.
329
+ self.visual_boundary_reentry = bool(visual_boundary_reentry)
330
+ # Source-centered boundary recurrence. The text-only L21 state is a
331
+ # fixed anchor, while the native multimodal residual initializes a
332
+ # nonzero, image-dependent deviation. The recurrent cell evolves
333
+ # only that deviation after subtracting its zero-deviation response.
334
+ self.visual_source_centered_reentry = bool(
335
+ visual_source_centered_reentry
336
+ )
337
+ # Restrict LoRA to the single shared recurrent cell L(ell_star+1).
338
+ # Its execution is disabled for R1, the text-only cache, answer-token
339
+ # decoding, and L22+; only recurrent transitions k>=2 enable it.
340
+ self.recurrence_only_adapter = bool(recurrence_only_adapter)
341
+ # Preserve the pretrained image-caused transition while removing the
342
+ # frozen layer's non-visual recurrent drift:
343
+ # Delta_k = F_{base+LoRA,on}(E,R_{k-1}) - F_{base,off}(E,R_{k-1})
344
+ # = (F_{base,on} - F_{base,off})
345
+ # + (F_{base+LoRA,on} - F_{base,on}).
346
+ # This is a full-state recurrence subtype and introduces no new
347
+ # parameters or inference-time module.
348
+ self.visual_preserving_adapter_correction = bool(
349
+ visual_preserving_adapter_correction
350
+ )
351
+ # Local Visual Latent Chain (LVLC): K homogeneous latent-token
352
+ # positions are evaluated in one native Transformer forward. The
353
+ # first latent sees clean Q plus visual source rows; every later latent
354
+ # sees only the immediately preceding latent plus the same visual
355
+ # rows; answer rows see clean Q plus the final latent. This is a local
356
+ # attention graph, not a recurrent forward or temporal aggregation.
357
+ self.local_visual_latent_chain = bool(local_visual_latent_chain)
358
+ # Perceive--Deliberate Latent Chain (PDLC). K must be even and is
359
+ # interpreted as K/2 native LOOK/THINK pairs. The complete VLM stack
360
+ # runs once over a block-sparse graph: LOOK rows may inspect the
361
+ # multimodal prompt, THINK rows may inspect only clean question rows
362
+ # and their paired LOOK, and answer rows may inspect only clean
363
+ # question rows and THINK states. There is no recurrent cell, visual
364
+ # state update, split-state splice, or temporal aggregator.
365
+ # LVLC reuses the same one-pass latent-sequence execution plumbing but
366
+ # replaces the LOOK/THINK token types and visibility graph entirely.
367
+ self.perceive_deliberate_chain = bool(
368
+ perceive_deliberate_chain or self.local_visual_latent_chain
369
+ )
370
+ # Tight PDLC graph: pair 1 bootstraps from prompt/question, later
371
+ # LOOK rows receive visual placeholders plus the preceding latent
372
+ # prefix, later THINK rows receive only that causal latent prefix, and
373
+ # the answer consumes final THINK. It changes visibility only and
374
+ # allocates no additional trainable module.
375
+ self.pdlc_transition_conditioned = bool(
376
+ pdlc_transition_conditioned
377
+ )
378
+ # Strict inference uses 1: only the bootstrap pair reads the clean
379
+ # question. Larger values are an explicit evaluation intervention.
380
+ # The training-only curriculum changes a runtime override and always
381
+ # returns to 1 for validation/inference.
382
+ self.pdlc_question_visible_through_pair = int(
383
+ pdlc_question_visible_through_pair
384
+ )
385
+ self.pdlc_question_curriculum = bool(pdlc_question_curriculum)
386
+ # Both learned latent token types start from one native text embedding
387
+ # (the launcher supplies the tokenizer's newline id). -1 falls back
388
+ # to token 0 for synthetic/unit-test configs only.
389
+ self.latent_chain_init_token_id = int(latent_chain_init_token_id)
390
+ # Spatial Recurrent Visual Field (SRVF). Unlike all historical
391
+ # ell_star methods, the recurrent state is the complete native visual
392
+ # token grid immediately after the vision merger. The final field is
393
+ # inserted into the ordinary Qwen prefix and the full decoder runs from
394
+ # scratch, making the zero-output initialization exactly base-equivalent.
395
+ self.spatial_visual_recurrence = bool(spatial_visual_recurrence)
396
+ self.spatial_recurrence_steps = int(spatial_recurrence_steps)
397
+ self.spatial_recurrence_inner_width = int(
398
+ spatial_recurrence_inner_width
399
+ )
400
+ self.spatial_recurrence_heads = int(spatial_recurrence_heads)
401
+ # -1 uses the trained/default horizon. Non-negative values are an
402
+ # evaluation-only prefix sweep, with 0 being the exact base path.
403
+ self.spatial_step_override = int(spatial_step_override)
404
+ self.counterfactual_beta = float(counterfactual_beta)
405
+ # Learn one feature-wise mixture of the complete recurrent residual
406
+ # trajectory. The bias-free T*d -> d map starts from
407
+ # [I/T, ..., I/T], so enabling it is exactly the historical CRR mean
408
+ # interface before the first optimizer step. A dedicated flag keeps
409
+ # old CRR checkpoints and the historical slot-ER aggregator distinct.
410
+ self.crr_concat_aggregation = bool(crr_concat_aggregation)
411
+ # Optional latent-only policy adapter. It corrects the deterministic
412
+ # C2..CT means but never touches B, C1, or answer-token decoding. Rank 0
413
+ # is the exact historical CRR path; RL checkpoints serialize rank > 0.
414
+ self.crr_policy_rank = int(crr_policy_rank)
415
+ self.crr_policy_alpha = float(crr_policy_alpha)
416
+ # Training-only CRR intervention. A batch-matched residual from a
417
+ # different-answer example replaces the factual residual before the
418
+ # upper decoder. The model returns the factual-minus-swapped answer
419
+ # score gap; the Trainer applies the margin loss. Evaluation and
420
+ # generation never execute the extra branch.
421
+ self.counterfactual_reliance_loss = bool(
422
+ counterfactual_reliance_loss
423
+ )
424
+ # Replace the aggregate residual swap with one intervention inside the
425
+ # recurrent chain. C_k is swapped and C_{k+1:T} is rolled out again,
426
+ # so the answer gap is attributable to a selected transition rather
427
+ # than to an undifferentiated full-trajectory replacement.
428
+ self.transition_aware_reliance_loss = bool(
429
+ transition_aware_reliance_loss
430
+ )
431
+ # Training-only hard-case corrective ranking. At one selected
432
+ # transition k>=2, the actual decoder scores Z_{k-1}, Z_k, and a
433
+ # donor-swapped Z_k. The Trainer requires the factual transition to
434
+ # improve over both detached comparators. This replaces, rather than
435
+ # stacks with, the historical transition-aware reliance objective.
436
+ self.corrective_transition_reliance_loss = bool(
437
+ corrective_transition_reliance_loss
438
+ )
439
+ # Training-only functional no-regression objective. Gold-answer NLL
440
+ # is measured through the actual upper decoder at every recurrent CRR
441
+ # prefix; inference remains the unchanged T-step mean interface.
442
+ self.counterfactual_step_retention_loss = bool(
443
+ counterfactual_step_retention_loss
444
+ )
445
+ # Training-only self-distillation target. The multimodal lower pass
446
+ # already computed for V* also provides its image-conditioned question
447
+ # rows. They supervise R_agg at the split boundary, but are never
448
+ # exposed to the inference decoder.
449
+ self.functional_state_loss = bool(functional_state_loss)
450
+ # Token CE lets long caption answers dominate one-token visual-choice
451
+ # examples. The per-example variant first averages valid answer tokens
452
+ # within each sample, then averages samples. It changes training/eval
453
+ # loss reduction only and has no inference-time effect.
454
+ self.per_example_answer_loss = bool(per_example_answer_loss)
455
+ # ER v2 (advisor 2026-08-11): explicit evidence read —
456
+ # R^(k) = f_theta(R^(k-1), Q*; E), E는 별도 memory로 cross-attn 읽기.
457
+ # (E+R 합산-감산 제거. evidence_reasoning=True 전제.)
458
+ self.explicit_evidence_read = bool(kwargs.pop("explicit_evidence_read", False))
459
+ # Historical field name retained for checkpoint compatibility. In ER
460
+ # v2.2 this means R^(0) is a batch-broadcast set of learnable slots; it
461
+ # is deliberately NOT initialized from Q*. Q* conditions every shared
462
+ # transition instead.
463
+ self.r_init_question = bool(kwargs.pop("r_init_question", False))
464
+ # ER v2.1 (advisor 2026-08-11): evidence read의 competitive normalization —
465
+ # slot 축 softmax 후 evidence 축 재정규화 (Slot Attention식 경쟁).
466
+ # 파라미터/모듈/loss 무추가; evidence_read의 W_qkv/W_o 재사용.
467
+ self.competitive_evidence_read = bool(
468
+ kwargs.pop("competitive_evidence_read", False))
469
+ # ER v2.2 C: k>=1 traj 감독 부재 시 twin 체인 제거 (states = answer 체인).
470
+ self.er_single_chain = bool(kwargs.pop("er_single_chain", False))
471
+ # ER v2.2 C (§5-7): decoder latent = softmax(temporal_logits)로 가중한
472
+ # R1..RT 결합 (전역 스칼라 T개). er_single_chain 전제. (legacy)
473
+ self.temporal_aggregation = bool(kwargs.pop("temporal_aggregation", False))
474
+ # ER v2.2 C 확정판: R_agg = Linear(Concat(R1..RT), bias=False),
475
+ # 초기 W=[I/T,...,I/T] (R_agg=mean(R1..RT)). er_single_chain 전제.
476
+ self.concat_aggregation = bool(kwargs.pop("concat_aggregation", False))
477
+ # ER v2: 워크스페이스 블록의 rank-r 부분공간 attention/FFN (0 = dense).
478
+ # d_w=d에서 state width 유지와 파라미터 폭증을 분리 (advisor).
479
+ self.workspace_low_rank = int(kwargs.pop("workspace_low_rank", 0))
480
+ # advisor: single-layer text-side interface alignment (bridge-only).
481
+ # teacher = frozen-base mm 경로 q_end@l*+1 (adapter off, sg);
482
+ # student = detach(R^T) splice의 l*+1 한 층 재계산. 추론 시 무존재.
483
+ self.interface_loss = interface_loss
484
+ # Training-only visual anchor switch. False in the current answer-only
485
+ # method; when false the model does not even materialize pooled V* for
486
+ # the Trainer.
487
+ self.read_anchor = bool(read_anchor)
488
+ if evidence_reasoning and persistent_slot_id:
489
+ raise ValueError("evidence_reasoning은 persistent_slot_id와 동시 사용 불가")
490
+ if evidence_reasoning and not mid_read_step:
491
+ raise ValueError("evidence_reasoning은 mid_read_step > 0 필요")
492
+ self.validate()
493
+
494
+ # -- derived views -----------------------------------------------------
495
+
496
+ @property
497
+ def num_decoder_layers(self) -> int:
498
+ return self.text_config.num_hidden_layers
499
+
500
+ @property
501
+ def backbone_width(self) -> int:
502
+ """``d`` -- backbone hidden width (3584 for the 7B)."""
503
+ return self.text_config.hidden_size
504
+
505
+ @property
506
+ def lower_slice(self) -> slice:
507
+ """Layers forming ``F_{<=l*}``."""
508
+ self._require_ell_star()
509
+ return slice(0, self.ell_star + 1)
510
+
511
+ @property
512
+ def upper_slice(self) -> slice:
513
+ """Layers forming ``F_{>l*}``."""
514
+ self._require_ell_star()
515
+ return slice(self.ell_star + 1, self.num_decoder_layers)
516
+
517
+ # -- validation --------------------------------------------------------
518
+
519
+ def _require_ell_star(self) -> None:
520
+ if self.ell_star is None:
521
+ raise ValueError(
522
+ "ell_star is unset. Run scripts/localize_read_point.py (Method 3.1) "
523
+ "and pass the measured layer explicitly; it must not be guessed."
524
+ )
525
+
526
+ def validate(self) -> None:
527
+ """Reject configurations that silently break the §3.2 contract."""
528
+ # ``PretrainedConfig.from_pretrained`` may apply user overrides with
529
+ # setattr *after* __init__. In that path, setting LVLC=True would not
530
+ # re-run the constructor's promotion of the shared one-pass execution
531
+ # flag. Re-establish the invariant before any model reads the config.
532
+ if getattr(self, "local_visual_latent_chain", False):
533
+ self.perceive_deliberate_chain = True
534
+ if self.ell_star is not None:
535
+ n = self.num_decoder_layers
536
+ # Both branches must be non-empty: an empty lower branch means there
537
+ # is no V* to read, an empty upper branch means the workspace never
538
+ # reaches the decoder.
539
+ if not 0 <= self.ell_star <= n - 2:
540
+ raise ValueError(
541
+ f"ell_star={self.ell_star} out of range for {n} decoder layers; "
542
+ f"expected 0 <= ell_star <= {n - 2} so both branches are non-empty."
543
+ )
544
+ # K=0 is the no-recurrence control (read only), not a misconfiguration.
545
+ if self.num_workspace_steps < 0:
546
+ raise ValueError("num_workspace_steps (K) must be >= 0.")
547
+ if self.num_workspace_slots < 1:
548
+ raise ValueError("num_workspace_slots (S) must be >= 1.")
549
+ if self.workspace_num_heads < 1:
550
+ raise ValueError("workspace_num_heads must be >= 1.")
551
+ if self.workspace_width % self.workspace_num_heads != 0:
552
+ raise ValueError(
553
+ f"workspace_width={self.workspace_width} must be divisible by "
554
+ f"workspace_num_heads={self.workspace_num_heads}."
555
+ )
556
+ if self.workspace_width > self.backbone_width:
557
+ raise ValueError(
558
+ f"workspace_width (d_w={self.workspace_width}) cannot exceed "
559
+ f"backbone width (d={self.backbone_width}). Native-width d_w=d is "
560
+ "the current method and uses an Identity decoder interface."
561
+ )
562
+ if self.workspace_rope_mode not in WORKSPACE_ROPE_MODES:
563
+ raise ValueError(
564
+ f"workspace_rope_mode={self.workspace_rope_mode!r} not in "
565
+ f"{WORKSPACE_ROPE_MODES}."
566
+ )
567
+ if self.adapter_rank < 0:
568
+ raise ValueError("adapter_rank must be >= 0 (0 disables the adapter).")
569
+ if not math.isfinite(self.adapter_dropout) or not 0.0 <= self.adapter_dropout < 1.0:
570
+ raise ValueError("adapter_dropout must be finite and in [0, 1).")
571
+ if getattr(self, "crr_policy_rank", 0) < 0:
572
+ raise ValueError("crr_policy_rank must be >= 0.")
573
+ if not math.isfinite(getattr(self, "crr_policy_alpha", 0.0)) or getattr(
574
+ self, "crr_policy_alpha", 0.0
575
+ ) <= 0.0:
576
+ raise ValueError("crr_policy_alpha must be finite and > 0.")
577
+ if self.read_num_blocks < 1 or self.transition_num_blocks < 1:
578
+ raise ValueError("read_num_blocks and transition_num_blocks must be >= 1.")
579
+ if self.workspace_ffn_mult < 1:
580
+ raise ValueError("workspace_ffn_mult must be >= 1.")
581
+ mid = int(getattr(self, "mid_read_step", 0) or 0)
582
+ if mid < 0 or mid > self.num_workspace_steps:
583
+ raise ValueError(
584
+ f"mid_read_step={mid} must be in [0, num_workspace_steps="
585
+ f"{self.num_workspace_steps}]."
586
+ )
587
+ low_rank = int(getattr(self, "workspace_low_rank", 0) or 0)
588
+ if low_rank < 0:
589
+ raise ValueError("workspace_low_rank must be >= 0.")
590
+ if low_rank and low_rank % self.workspace_num_heads != 0:
591
+ raise ValueError(
592
+ f"workspace_low_rank={low_rank} must be divisible by "
593
+ f"workspace_num_heads={self.workspace_num_heads}."
594
+ )
595
+ if getattr(self, "competitive_evidence_read", False) and not getattr(
596
+ self, "explicit_evidence_read", False
597
+ ):
598
+ raise ValueError(
599
+ "competitive_evidence_read requires explicit_evidence_read=True; "
600
+ "otherwise the flag has no execution path."
601
+ )
602
+ if getattr(self, "explicit_evidence_read", False) and not getattr(
603
+ self, "evidence_reasoning", False
604
+ ):
605
+ raise ValueError(
606
+ "explicit_evidence_read requires evidence_reasoning=True; otherwise "
607
+ "the evidence-read module is allocated but never called."
608
+ )
609
+ if getattr(self, "r_init_question", False) and not getattr(
610
+ self, "evidence_reasoning", False
611
+ ):
612
+ raise ValueError(
613
+ "r_init_question (legacy name for learnable R0 slots) requires "
614
+ "evidence_reasoning=True; otherwise reasoning_slots are unused."
615
+ )
616
+ if (getattr(self, "concat_aggregation", False)
617
+ or getattr(self, "temporal_aggregation", False)):
618
+ if self.num_workspace_steps < 1:
619
+ raise ValueError("reasoning aggregation requires num_workspace_steps >= 1.")
620
+ if not getattr(self, "evidence_reasoning", False):
621
+ raise ValueError(
622
+ "reasoning aggregation requires evidence_reasoning=True; "
623
+ "otherwise the aggregation module is unused."
624
+ )
625
+ if getattr(self, "concat_aggregation", False) and not getattr(
626
+ self, "er_single_chain", False
627
+ ):
628
+ raise ValueError("concat_aggregation requires er_single_chain=True.")
629
+ if getattr(self, "raw_evidence_question", False):
630
+ if self.workspace_width != self.backbone_width:
631
+ raise ValueError(
632
+ "raw_evidence_question requires native workspace_width == "
633
+ "backbone_width: E=V* and R0=Q* are used without projections."
634
+ )
635
+ if self.num_workspace_steps < 1:
636
+ raise ValueError(
637
+ "raw_evidence_question requires num_workspace_steps >= 1."
638
+ )
639
+ source_layer = (
640
+ self.ell_star + 1
641
+ if self.raw_attention_source_layer < 0
642
+ else self.raw_attention_source_layer
643
+ )
644
+ if not self.ell_star < source_layer < self.text_config.num_hidden_layers:
645
+ raise ValueError(
646
+ "raw_attention_source_layer must be an upper-backbone layer, "
647
+ f"got ell_star={self.ell_star}, source={source_layer}."
648
+ )
649
+ incompatible = {
650
+ "evidence_reasoning": bool(getattr(self, "evidence_reasoning", False)),
651
+ "explicit_evidence_read": bool(getattr(self, "explicit_evidence_read", False)),
652
+ "competitive_evidence_read": bool(getattr(self, "competitive_evidence_read", False)),
653
+ "r_init_question": bool(getattr(self, "r_init_question", False)),
654
+ "er_single_chain": bool(getattr(self, "er_single_chain", False)),
655
+ "temporal_aggregation": bool(getattr(self, "temporal_aggregation", False)),
656
+ "concat_aggregation": bool(getattr(self, "concat_aggregation", False)),
657
+ "read_shortcut": bool(getattr(self, "read_shortcut", False)),
658
+ "splice_scale_match": bool(getattr(self, "splice_scale_match", False)),
659
+ "read_gating": bool(getattr(self, "read_gating", False)),
660
+ "read_residual_pure": bool(getattr(self, "read_residual_pure", False)),
661
+ "final_verify": bool(getattr(self, "final_verify", False)),
662
+ "mid_read_step": bool(getattr(self, "mid_read_step", 0)),
663
+ "persistent_slot_id": bool(getattr(self, "persistent_slot_id", False)),
664
+ "interface_loss": bool(getattr(self, "interface_loss", False)),
665
+ "read_anchor": bool(getattr(self, "read_anchor", False)),
666
+ }
667
+ active = [name for name, enabled in incompatible.items() if enabled]
668
+ if active:
669
+ raise ValueError(
670
+ "raw_evidence_question replaces the slot workspace and is "
671
+ f"incompatible with: {', '.join(active)}"
672
+ )
673
+ if self.raw_state_replacement and self.raw_transition_ffn:
674
+ raise ValueError(
675
+ "raw_state_replacement is CrossAttn-only and is incompatible "
676
+ "with raw_transition_ffn=True."
677
+ )
678
+ if self.raw_question_anchor and not self.raw_state_replacement:
679
+ raise ValueError(
680
+ "raw_question_anchor requires raw_state_replacement=True."
681
+ )
682
+ if self.raw_mean_aggregation and (
683
+ self.raw_state_replacement or self.raw_question_anchor
684
+ ):
685
+ raise ValueError(
686
+ "raw_mean_aggregation replaces raw_state_replacement and "
687
+ "raw_question_anchor; use the residual recurrence."
688
+ )
689
+ hierarchy = list(
690
+ getattr(self, "hierarchical_visual_layers", ()) or ()
691
+ )
692
+ if hierarchy:
693
+ if not self.raw_mean_aggregation:
694
+ raise ValueError(
695
+ "hierarchical_visual_layers requires "
696
+ "raw_mean_aggregation=True."
697
+ )
698
+ if self.raw_transition_ffn:
699
+ raise ValueError(
700
+ "hierarchical visual recurrence is CrossAttn-only and "
701
+ "requires raw_transition_ffn=False."
702
+ )
703
+ expected_visual_reads = self.num_workspace_steps - 1
704
+ if len(hierarchy) != expected_visual_reads:
705
+ raise ValueError(
706
+ "hierarchical_visual_layers must contain exactly one "
707
+ "layer per visual-read step, followed by one Q* "
708
+ "integration step: "
709
+ f"got {len(hierarchy)} layers for "
710
+ f"T={self.num_workspace_steps} "
711
+ f"(expected {expected_visual_reads})."
712
+ )
713
+ if hierarchy != sorted(set(hierarchy)):
714
+ raise ValueError(
715
+ "hierarchical_visual_layers must be strictly increasing."
716
+ )
717
+ invalid = [
718
+ layer
719
+ for layer in hierarchy
720
+ if not 0 <= layer <= int(self.ell_star)
721
+ ]
722
+ if invalid:
723
+ raise ValueError(
724
+ "hierarchical_visual_layers must lie in the frozen lower "
725
+ f"branch [0,{self.ell_star}], got {invalid}."
726
+ )
727
+ if self.functional_state_loss:
728
+ raise ValueError(
729
+ "functional_state_loss is not defined for hierarchical "
730
+ "visual recurrence; use final-answer CE only."
731
+ )
732
+ if self.functional_state_loss and not self.raw_mean_aggregation:
733
+ raise ValueError(
734
+ "functional_state_loss requires raw_mean_aggregation=True "
735
+ "so its target is the exact decoder-facing recurrence."
736
+ )
737
+ elif getattr(self, "raw_state_replacement", False):
738
+ raise ValueError(
739
+ "raw_state_replacement requires raw_evidence_question=True."
740
+ )
741
+ elif getattr(self, "raw_question_anchor", False):
742
+ raise ValueError(
743
+ "raw_question_anchor requires raw_evidence_question=True."
744
+ )
745
+ elif getattr(self, "raw_mean_aggregation", False):
746
+ raise ValueError(
747
+ "raw_mean_aggregation requires raw_evidence_question=True."
748
+ )
749
+ elif getattr(self, "hierarchical_visual_layers", ()):
750
+ raise ValueError(
751
+ "hierarchical_visual_layers requires "
752
+ "raw_evidence_question=True."
753
+ )
754
+ elif getattr(self, "functional_state_loss", False):
755
+ raise ValueError(
756
+ "functional_state_loss requires raw_evidence_question=True."
757
+ )
758
+ if (
759
+ getattr(self, "counterfactual_reliance_loss", False)
760
+ and not getattr(self, "counterfactual_residual_recurrence", False)
761
+ ):
762
+ raise ValueError(
763
+ "counterfactual_reliance_loss requires "
764
+ "counterfactual_residual_recurrence=True."
765
+ )
766
+ if (
767
+ getattr(self, "transition_aware_reliance_loss", False)
768
+ and not getattr(self, "counterfactual_reliance_loss", False)
769
+ ):
770
+ raise ValueError(
771
+ "transition_aware_reliance_loss requires "
772
+ "counterfactual_reliance_loss=True."
773
+ )
774
+ if (
775
+ getattr(self, "transition_aware_reliance_loss", False)
776
+ and self.num_workspace_steps < 2
777
+ ):
778
+ raise ValueError(
779
+ "transition_aware_reliance_loss requires at least two steps."
780
+ )
781
+ if (
782
+ getattr(self, "corrective_transition_reliance_loss", False)
783
+ and not getattr(self, "counterfactual_reliance_loss", False)
784
+ ):
785
+ raise ValueError(
786
+ "corrective_transition_reliance_loss requires "
787
+ "counterfactual_reliance_loss=True."
788
+ )
789
+ if (
790
+ getattr(self, "corrective_transition_reliance_loss", False)
791
+ and self.num_workspace_steps < 2
792
+ ):
793
+ raise ValueError(
794
+ "corrective_transition_reliance_loss requires at least two steps."
795
+ )
796
+ if (
797
+ getattr(self, "corrective_transition_reliance_loss", False)
798
+ and getattr(self, "transition_aware_reliance_loss", False)
799
+ ):
800
+ raise ValueError(
801
+ "corrective_transition_reliance_loss replaces, and cannot be "
802
+ "combined with, transition_aware_reliance_loss."
803
+ )
804
+ if (
805
+ getattr(self, "corrective_transition_reliance_loss", False)
806
+ and getattr(self, "counterfactual_step_retention_loss", False)
807
+ ):
808
+ raise ValueError(
809
+ "corrective_transition_reliance_loss replaces, and cannot be "
810
+ "combined with, counterfactual_step_retention_loss."
811
+ )
812
+ if (
813
+ getattr(self, "counterfactual_step_retention_loss", False)
814
+ and not getattr(self, "counterfactual_residual_recurrence", False)
815
+ ):
816
+ raise ValueError(
817
+ "counterfactual_step_retention_loss requires "
818
+ "counterfactual_residual_recurrence=True."
819
+ )
820
+ if (
821
+ getattr(self, "counterfactual_step_retention_loss", False)
822
+ and self.num_workspace_steps < 2
823
+ ):
824
+ raise ValueError(
825
+ "counterfactual_step_retention_loss requires at least two steps."
826
+ )
827
+ if getattr(self, "crr_concat_aggregation", False) and not getattr(
828
+ self, "counterfactual_residual_recurrence", False
829
+ ):
830
+ raise ValueError(
831
+ "crr_concat_aggregation requires "
832
+ "counterfactual_residual_recurrence=True."
833
+ )
834
+ if getattr(self, "counterfactual_residual_recurrence", False):
835
+ if getattr(self, "raw_evidence_question", False):
836
+ raise ValueError(
837
+ "counterfactual_residual_recurrence replaces "
838
+ "raw_evidence_question; enable exactly one q-state method."
839
+ )
840
+ if self.num_workspace_steps < 1:
841
+ raise ValueError(
842
+ "counterfactual_residual_recurrence requires "
843
+ "num_workspace_steps >= 1."
844
+ )
845
+ if self.crr_policy_rank > 0 and self.num_workspace_steps < 2:
846
+ raise ValueError(
847
+ "crr_policy_rank > 0 requires num_workspace_steps >= 2."
848
+ )
849
+ if not 0.0 <= self.counterfactual_beta <= 1.0:
850
+ raise ValueError(
851
+ "counterfactual_beta must lie in [0, 1], got "
852
+ f"{self.counterfactual_beta}."
853
+ )
854
+ incompatible = {
855
+ "evidence_reasoning": bool(getattr(self, "evidence_reasoning", False)),
856
+ "explicit_evidence_read": bool(getattr(self, "explicit_evidence_read", False)),
857
+ "competitive_evidence_read": bool(getattr(self, "competitive_evidence_read", False)),
858
+ "r_init_question": bool(getattr(self, "r_init_question", False)),
859
+ "er_single_chain": bool(getattr(self, "er_single_chain", False)),
860
+ "temporal_aggregation": bool(getattr(self, "temporal_aggregation", False)),
861
+ "concat_aggregation": bool(getattr(self, "concat_aggregation", False)),
862
+ "read_shortcut": bool(getattr(self, "read_shortcut", False)),
863
+ "splice_scale_match": bool(getattr(self, "splice_scale_match", False)),
864
+ "read_gating": bool(getattr(self, "read_gating", False)),
865
+ "read_residual_pure": bool(getattr(self, "read_residual_pure", False)),
866
+ "final_verify": bool(getattr(self, "final_verify", False)),
867
+ "mid_read_step": bool(getattr(self, "mid_read_step", 0)),
868
+ "persistent_slot_id": bool(getattr(self, "persistent_slot_id", False)),
869
+ "interface_loss": bool(getattr(self, "interface_loss", False)),
870
+ "read_anchor": bool(getattr(self, "read_anchor", False)),
871
+ "functional_state_loss": bool(getattr(self, "functional_state_loss", False)),
872
+ }
873
+ active = [name for name, enabled in incompatible.items() if enabled]
874
+ if active:
875
+ raise ValueError(
876
+ "counterfactual_residual_recurrence replaces the slot/raw "
877
+ f"workspace and is incompatible with: {', '.join(active)}"
878
+ )
879
+ elif getattr(self, "crr_policy_rank", 0) > 0:
880
+ raise ValueError(
881
+ "crr_policy_rank > 0 requires "
882
+ "counterfactual_residual_recurrence=True."
883
+ )
884
+ if getattr(self, "visual_counterfactual_recurrence", False):
885
+ if not getattr(
886
+ self, "counterfactual_residual_recurrence", False
887
+ ):
888
+ raise ValueError(
889
+ "visual_counterfactual_recurrence requires "
890
+ "counterfactual_residual_recurrence=True."
891
+ )
892
+ # CVRR is the intentionally minimal CE-only candidate. The old
893
+ # CRR tail-rollout losses/policy implement a different transition
894
+ # equation and must not be silently applied to this trajectory.
895
+ incompatible = {
896
+ "counterfactual_reliance_loss": bool(
897
+ getattr(self, "counterfactual_reliance_loss", False)
898
+ ),
899
+ "transition_aware_reliance_loss": bool(
900
+ getattr(self, "transition_aware_reliance_loss", False)
901
+ ),
902
+ "corrective_transition_reliance_loss": bool(
903
+ getattr(self, "corrective_transition_reliance_loss", False)
904
+ ),
905
+ "counterfactual_step_retention_loss": bool(
906
+ getattr(self, "counterfactual_step_retention_loss", False)
907
+ ),
908
+ "crr_concat_aggregation": bool(
909
+ getattr(self, "crr_concat_aggregation", False)
910
+ ),
911
+ "crr_policy_rank": bool(getattr(self, "crr_policy_rank", 0)),
912
+ }
913
+ active = [name for name, enabled in incompatible.items() if enabled]
914
+ if active:
915
+ raise ValueError(
916
+ "visual_counterfactual_recurrence uses final-answer CE only "
917
+ "and a parameter-free decoder interface; incompatible with: "
918
+ + ", ".join(active)
919
+ )
920
+ elif getattr(self, "visual_cumulative_recurrence", False):
921
+ raise ValueError(
922
+ "visual_cumulative_recurrence requires "
923
+ "visual_counterfactual_recurrence=True."
924
+ )
925
+ if getattr(self, "visual_full_state_recurrence", False):
926
+ if not getattr(self, "visual_counterfactual_recurrence", False):
927
+ raise ValueError(
928
+ "visual_full_state_recurrence requires "
929
+ "visual_counterfactual_recurrence=True."
930
+ )
931
+ if not getattr(self, "visual_cumulative_recurrence", False):
932
+ raise ValueError(
933
+ "visual_full_state_recurrence requires "
934
+ "visual_cumulative_recurrence=True."
935
+ )
936
+ if getattr(self, "visual_boundary_reentry", False) and not getattr(
937
+ self, "visual_full_state_recurrence", False
938
+ ):
939
+ raise ValueError(
940
+ "visual_boundary_reentry requires "
941
+ "visual_full_state_recurrence=True."
942
+ )
943
+ if getattr(self, "visual_source_centered_reentry", False):
944
+ if not getattr(self, "visual_boundary_reentry", False):
945
+ raise ValueError(
946
+ "visual_source_centered_reentry requires "
947
+ "visual_boundary_reentry=True."
948
+ )
949
+ if getattr(self, "visual_preserving_adapter_correction", False):
950
+ raise ValueError(
951
+ "visual_source_centered_reentry cannot be combined with "
952
+ "visual_preserving_adapter_correction."
953
+ )
954
+ if getattr(self, "recurrence_only_adapter", False):
955
+ if not getattr(self, "visual_full_state_recurrence", False):
956
+ raise ValueError(
957
+ "recurrence_only_adapter requires "
958
+ "visual_full_state_recurrence=True."
959
+ )
960
+ if self.adapter_rank <= 0:
961
+ raise ValueError(
962
+ "recurrence_only_adapter requires adapter_rank > 0."
963
+ )
964
+ if self.ell_star is None:
965
+ raise ValueError(
966
+ "recurrence_only_adapter requires an explicit ell_star."
967
+ )
968
+ recurrent_layer = int(self.ell_star) + 1
969
+ if recurrent_layer in set(self.adapter_exclude_layers):
970
+ raise ValueError(
971
+ "recurrence_only_adapter cannot exclude its recurrent "
972
+ f"layer {recurrent_layer} from LoRA."
973
+ )
974
+ if getattr(self, "visual_preserving_adapter_correction", False):
975
+ if not getattr(self, "visual_full_state_recurrence", False):
976
+ raise ValueError(
977
+ "visual_preserving_adapter_correction requires "
978
+ "visual_full_state_recurrence=True."
979
+ )
980
+ if not getattr(self, "recurrence_only_adapter", False):
981
+ raise ValueError(
982
+ "visual_preserving_adapter_correction requires "
983
+ "recurrence_only_adapter=True."
984
+ )
985
+ if getattr(self, "perceive_deliberate_chain", False):
986
+ if self.num_workspace_steps < 2 or self.num_workspace_steps % 2:
987
+ raise ValueError(
988
+ "one-pass latent chains require an even "
989
+ "num_workspace_steps >= 2."
990
+ )
991
+ vocab_size = int(self.text_config.vocab_size)
992
+ if not -1 <= self.latent_chain_init_token_id < vocab_size:
993
+ raise ValueError(
994
+ "latent_chain_init_token_id must be -1 or a valid text "
995
+ f"token id below {vocab_size}."
996
+ )
997
+ num_pairs = self.num_workspace_steps // 2
998
+ if self.local_visual_latent_chain and (
999
+ self.pdlc_transition_conditioned
1000
+ or self.pdlc_question_curriculum
1001
+ or self.pdlc_question_visible_through_pair != 1
1002
+ ):
1003
+ raise ValueError(
1004
+ "local_visual_latent_chain replaces PDLC transition and "
1005
+ "question-visibility controls"
1006
+ )
1007
+ if not 1 <= self.pdlc_question_visible_through_pair <= num_pairs:
1008
+ raise ValueError(
1009
+ "pdlc_question_visible_through_pair must lie in "
1010
+ f"[1,{num_pairs}]"
1011
+ )
1012
+ if self.pdlc_question_curriculum and not (
1013
+ getattr(self, "pdlc_transition_conditioned", False)
1014
+ ):
1015
+ raise ValueError(
1016
+ "pdlc_question_curriculum requires "
1017
+ "pdlc_transition_conditioned=True"
1018
+ )
1019
+ incompatible = {
1020
+ "raw_evidence_question": bool(
1021
+ getattr(self, "raw_evidence_question", False)
1022
+ ),
1023
+ "counterfactual_residual_recurrence": bool(
1024
+ getattr(self, "counterfactual_residual_recurrence", False)
1025
+ ),
1026
+ "spatial_visual_recurrence": bool(
1027
+ getattr(self, "spatial_visual_recurrence", False)
1028
+ ),
1029
+ "evidence_reasoning": bool(
1030
+ getattr(self, "evidence_reasoning", False)
1031
+ ),
1032
+ "read_gating": bool(getattr(self, "read_gating", False)),
1033
+ "read_residual_pure": bool(
1034
+ getattr(self, "read_residual_pure", False)
1035
+ ),
1036
+ "final_verify": bool(getattr(self, "final_verify", False)),
1037
+ "mid_read_step": bool(getattr(self, "mid_read_step", 0)),
1038
+ "interface_loss": bool(getattr(self, "interface_loss", False)),
1039
+ "read_anchor": bool(getattr(self, "read_anchor", False)),
1040
+ "counterfactual_reliance_loss": bool(
1041
+ getattr(self, "counterfactual_reliance_loss", False)
1042
+ ),
1043
+ "counterfactual_step_retention_loss": bool(
1044
+ getattr(self, "counterfactual_step_retention_loss", False)
1045
+ ),
1046
+ "functional_state_loss": bool(
1047
+ getattr(self, "functional_state_loss", False)
1048
+ ),
1049
+ }
1050
+ active = [name for name, enabled in incompatible.items() if enabled]
1051
+ if active:
1052
+ raise ValueError(
1053
+ "perceive_deliberate_chain replaces every recurrent/"
1054
+ "workspace path; incompatible with: " + ", ".join(active)
1055
+ )
1056
+ elif getattr(self, "pdlc_transition_conditioned", False):
1057
+ raise ValueError(
1058
+ "pdlc_transition_conditioned requires "
1059
+ "perceive_deliberate_chain=True."
1060
+ )
1061
+ elif getattr(self, "pdlc_question_curriculum", False) or (
1062
+ getattr(self, "pdlc_question_visible_through_pair", 1) != 1
1063
+ ):
1064
+ raise ValueError(
1065
+ "PDLC question visibility controls require "
1066
+ "perceive_deliberate_chain=True."
1067
+ )
1068
+ if getattr(self, "spatial_visual_recurrence", False):
1069
+ if not 0.0 <= self.counterfactual_beta <= 1.0:
1070
+ raise ValueError(
1071
+ "counterfactual_beta must lie in [0, 1], got "
1072
+ f"{self.counterfactual_beta}."
1073
+ )
1074
+ if self.spatial_recurrence_steps < 1:
1075
+ raise ValueError(
1076
+ "spatial_recurrence_steps must be >= 1 for training"
1077
+ )
1078
+ if self.spatial_recurrence_inner_width < 1:
1079
+ raise ValueError(
1080
+ "spatial_recurrence_inner_width must be positive"
1081
+ )
1082
+ if self.spatial_recurrence_heads < 1 or (
1083
+ self.spatial_recurrence_inner_width
1084
+ % self.spatial_recurrence_heads
1085
+ ):
1086
+ raise ValueError(
1087
+ "spatial_recurrence_inner_width must be divisible by "
1088
+ "spatial_recurrence_heads"
1089
+ )
1090
+ if self.spatial_step_override < -1:
1091
+ raise ValueError("spatial_step_override must be >= -1")
1092
+ if self.adapter_rank != 0:
1093
+ raise ValueError(
1094
+ "spatial_visual_recurrence replaces decoder LoRA; set "
1095
+ "adapter_rank=0 so the T=0 path remains the frozen base model"
1096
+ )
1097
+ incompatible = {
1098
+ "raw_evidence_question": bool(
1099
+ getattr(self, "raw_evidence_question", False)
1100
+ ),
1101
+ "counterfactual_residual_recurrence": bool(
1102
+ getattr(self, "counterfactual_residual_recurrence", False)
1103
+ ),
1104
+ "evidence_reasoning": bool(
1105
+ getattr(self, "evidence_reasoning", False)
1106
+ ),
1107
+ "read_gating": bool(getattr(self, "read_gating", False)),
1108
+ "read_residual_pure": bool(
1109
+ getattr(self, "read_residual_pure", False)
1110
+ ),
1111
+ "final_verify": bool(getattr(self, "final_verify", False)),
1112
+ "mid_read_step": bool(getattr(self, "mid_read_step", 0)),
1113
+ "interface_loss": bool(getattr(self, "interface_loss", False)),
1114
+ "read_anchor": bool(getattr(self, "read_anchor", False)),
1115
+ "counterfactual_reliance_loss": bool(
1116
+ getattr(self, "counterfactual_reliance_loss", False)
1117
+ ),
1118
+ "counterfactual_step_retention_loss": bool(
1119
+ getattr(self, "counterfactual_step_retention_loss", False)
1120
+ ),
1121
+ "functional_state_loss": bool(
1122
+ getattr(self, "functional_state_loss", False)
1123
+ ),
1124
+ }
1125
+ active = [name for name, enabled in incompatible.items() if enabled]
1126
+ if active:
1127
+ raise ValueError(
1128
+ "spatial_visual_recurrence replaces every historical "
1129
+ "workspace/reliance path; incompatible with: "
1130
+ + ", ".join(active)
1131
+ )
1132
+ modes = [bool(getattr(self, m, False)) for m in
1133
+ ("read_gating", "read_residual_pure", "final_verify")]
1134
+ modes.append(bool(getattr(self, "mid_read_step", 0)))
1135
+ if sum(modes) > 1:
1136
+ raise ValueError(
1137
+ "read_gating / read_residual_pure / final_verify are mutually "
1138
+ "exclusive reread modes."
1139
+ )
1140
+
1141
+
1142
+ __all__ = ["CloseQwen2_5_VLConfig", "WORKSPACE_ROPE_MODES"]
configuration_source_qwen3.py ADDED
@@ -0,0 +1,740 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Configuration for CLOSE latent reasoning on Qwen3-VL.
2
+
3
+ The current ``raw_evidence_question`` path keeps the complete layer-``ell_star``
4
+ visual memory as immutable evidence and recurrently updates the text-only
5
+ question state by cross-attending to it. Its selected interface averages all
6
+ four residual recurrent states and may use a training-only restoration target
7
+ from the normal multimodal question state. Historical slot-workspace fields are
8
+ retained so old checkpoints remain loadable, but they are mutually exclusive
9
+ with the raw-evidence path.
10
+
11
+ Every architectural hyperparameter of Method §3 lives here -- nothing is
12
+ hardcoded in ``modeling_close_qwen3_vl.py``.
13
+
14
+ Layer-index convention (used consistently across the repo)
15
+ ----------------------------------------------------------
16
+ ``ell_star`` is **0-indexed and inclusive**: it is the index of the last decoder
17
+ layer belonging to the lower branch. With ``num_hidden_layers == 28``::
18
+
19
+ F_{<=l*} = language_model.layers[: ell_star + 1] # lower_slice
20
+ F_{>l*} = language_model.layers[ell_star + 1 :] # upper_slice
21
+
22
+ so ``ell_star=18`` means layers 0..18 run before the split and 19..27 after it.
23
+ ``ell_star`` is measured once by ``scripts/localize_read_point.py`` (§3.1) and
24
+ frozen before any workspace training.
25
+
26
+ Careful with ``Qwen3VLConfig``: its ``__setattr__`` forwards any attribute
27
+ whose name already exists in ``text_config.__dict__`` down to the sub-config, and
28
+ its ``__init__`` builds ``text_config`` from ``**kwargs``. Workspace fields are
29
+ therefore declared as explicit named parameters so they never enter ``kwargs``
30
+ and never shadow a text-config field. ``tests/test_config_roundtrip.py`` guards
31
+ this.
32
+ """
33
+
34
+ from __future__ import annotations
35
+
36
+ import math
37
+
38
+ from transformers.models.qwen3_vl.configuration_qwen3_vl import Qwen3VLConfig
39
+
40
+ #: How the ``S`` workspace slots are assigned M-RoPE positions once they are
41
+ #: spliced into the replacement cache above ``ell_star``. Qwen3-VL uses 3D
42
+ #: M-RoPE (t, h, w) with ``mrope_section=[16, 24, 24]``, so slots -- which have
43
+ #: no spatial extent -- need an explicit convention.
44
+ WORKSPACE_ROPE_MODES = (
45
+ # Slots continue the 1-D text position sequence directly after Q*, with
46
+ # t == h == w. Treats the workspace as "more text".
47
+ "continue",
48
+ # Slots reuse the (t, h, w) span the image tokens occupied in the original
49
+ # multimodal prefill, subsampled to S positions. Preserves the read point's
50
+ # positional geometry but reintroduces image-derived indices.
51
+ "image_span",
52
+ # All slots share one position (the first position after Q*). Makes the
53
+ # workspace order-free, matching the set semantics of Perceiver-style slots.
54
+ "shared",
55
+ )
56
+
57
+
58
+ class CloseQwen3VLConfig(Qwen3VLConfig):
59
+ r"""Configuration for :class:`CloseQwen3VLForConditionalGeneration`.
60
+
61
+ Args:
62
+ ell_star: 0-indexed inclusive last layer of the lower branch. ``None``
63
+ until §3.1 localization has run; the model refuses to build with
64
+ ``None`` rather than falling back to a guessed depth.
65
+ num_workspace_steps: the fixed recurrent horizon (paper notation ``T``;
66
+ per-item segment counts ``n`` exist only in the losses, never here).
67
+ The trainable parameter count is independent of it because
68
+ ``f_theta`` is shared across steps, so it is a free experimental
69
+ axis. ``0`` is legal and is the no-recurrence floor: ``Z^(0)`` goes
70
+ straight to the decoder.
71
+ num_workspace_slots: ``S``, slots in ``Z``. Free capacity knob.
72
+ workspace_width: ``d_w``, the slot width. Free capacity knob -- the
73
+ method's claim is path exclusivity (all visual information routed
74
+ through ``Z``), NOT an information bottleneck, so ``S`` and ``d_w``
75
+ carry no methodological commitment. Kept below ``d`` (4096 for the
76
+ 7B) since ``P_Z`` projects up at the splice.
77
+ workspace_num_heads: attention heads inside ``r_theta`` / ``f_theta``.
78
+ read_num_blocks: cross-attention blocks in ``r_theta``. The *read* is
79
+ once regardless of depth: all blocks see the same ``V*`` in a single
80
+ forward, and ``V*`` is discarded afterwards.
81
+ transition_num_blocks: blocks in the shared recurrent ``f_theta``.
82
+ workspace_ffn_mult: FFN expansion inside workspace blocks.
83
+ workspace_dropout: dropout inside workspace blocks.
84
+ workspace_rope_mode: one of :data:`WORKSPACE_ROPE_MODES`.
85
+ adapter_rank: rank of the low-rank adapter on the post-``ell_star``
86
+ backbone layers. ``0`` disables it (frozen backbone only).
87
+ adapter_alpha: LoRA-style scaling; effective scale is
88
+ ``adapter_alpha / adapter_rank``.
89
+ adapter_dropout: dropout probability applied only to the LoRA branch.
90
+ lambda_traj: legacy trajectory-loss weight. The current method sets it
91
+ to zero and trains with final-answer CE only; it remains serialized
92
+ solely so historical checkpoints can still be inspected.
93
+ """
94
+
95
+ model_type = "close_qwen3_vl"
96
+
97
+ def __init__(
98
+ self,
99
+ ell_star: int | None = None,
100
+ num_workspace_steps: int = 4,
101
+ num_workspace_slots: int = 32,
102
+ workspace_width: int = 512,
103
+ workspace_num_heads: int = 8,
104
+ read_num_blocks: int = 1,
105
+ transition_num_blocks: int = 2,
106
+ workspace_ffn_mult: int = 4,
107
+ workspace_dropout: float = 0.0,
108
+ workspace_rope_mode: str = "continue",
109
+ adapter_rank: int = 16,
110
+ adapter_alpha: int = 32,
111
+ adapter_dropout: float = 0.0,
112
+ adapter_exclude_layers: tuple = (),
113
+ lambda_traj: float = 0.0,
114
+ read_shortcut: bool = False,
115
+ splice_scale_match: bool = False,
116
+ read_gating: bool = False,
117
+ gate_init_eps: float = 0.02,
118
+ read_residual_pure: bool = False,
119
+ residual_norm_cap: float = 0.0,
120
+ pure_warm_start: bool = False,
121
+ final_verify: bool = False,
122
+ mid_read_step: int = 0,
123
+ persistent_slot_id: bool = False,
124
+ mid_read_alpha: float = 1.0,
125
+ evidence_reasoning: bool = False,
126
+ raw_evidence_question: bool = False,
127
+ raw_attention_source_layer: int = -1,
128
+ raw_transition_ffn: bool | None = None,
129
+ raw_state_replacement: bool = False,
130
+ raw_question_anchor: bool = False,
131
+ raw_mean_aggregation: bool = False,
132
+ counterfactual_residual_recurrence: bool = False,
133
+ visual_counterfactual_recurrence: bool = False,
134
+ visual_cumulative_recurrence: bool = False,
135
+ visual_full_state_recurrence: bool = False,
136
+ recurrence_only_adapter: bool = False,
137
+ counterfactual_beta: float = 1.0,
138
+ crr_concat_aggregation: bool = False,
139
+ crr_policy_rank: int = 0,
140
+ crr_policy_alpha: float = 32.0,
141
+ counterfactual_reliance_loss: bool = False,
142
+ transition_aware_reliance_loss: bool = False,
143
+ counterfactual_step_retention_loss: bool = False,
144
+ functional_state_loss: bool = False,
145
+ per_example_answer_loss: bool = False,
146
+ interface_loss: bool = False,
147
+ read_anchor: bool = False,
148
+ **kwargs,
149
+ ):
150
+ super().__init__(**kwargs)
151
+ self.ell_star = ell_star
152
+ self.num_workspace_steps = num_workspace_steps
153
+ self.num_workspace_slots = num_workspace_slots
154
+ self.workspace_width = workspace_width
155
+ self.workspace_num_heads = workspace_num_heads
156
+ self.read_num_blocks = read_num_blocks
157
+ self.transition_num_blocks = transition_num_blocks
158
+ self.workspace_ffn_mult = workspace_ffn_mult
159
+ self.workspace_dropout = workspace_dropout
160
+ self.workspace_rope_mode = workspace_rope_mode
161
+ self.adapter_rank = adapter_rank
162
+ self.adapter_alpha = adapter_alpha
163
+ self.adapter_dropout = adapter_dropout
164
+ # E run: dense full-FT 층은 LoRA 대상에서 제외 (이중 파라미터화 방지)
165
+ self.adapter_exclude_layers = list(adapter_exclude_layers or ())
166
+ self.lambda_traj = lambda_traj
167
+ # v2 anti-collapse switches (defaults False so v1 checkpoints reload).
168
+ # read_shortcut: r_theta adds an orthogonally-projected pooled-V* term to
169
+ # every slot, so Z^(0) is image-dependent from step 0 instead of relying
170
+ # on a random cross-attention finding signal.
171
+ # splice_scale_match: P_Z(Z) rows are RMS-matched to the question rows at
172
+ # the splice. Unmatched, slot states are ~O(1) against layer-20 hiddens
173
+ # in the hundreds -- the decoder could ignore them for free, and did
174
+ # (probe 20143: visual sensitivity 0.94 -> 1.0000 over training).
175
+ self.read_shortcut = read_shortcut
176
+ self.splice_scale_match = splice_scale_match
177
+ # v6 candidate: gated re-reads DURING the recurrence. Every re-read
178
+ # goes through r_theta into Z, so the mediation invariant (decoder sees
179
+ # only [Q*; P_Z(Z)]) is untouched; what changes is that the workspace
180
+ # may consult V* again mid-reasoning, scaled by a learned scalar gate.
181
+ # The gate is free at EVERY step -- no hand-coded schedule. Degeneration
182
+ # into answer-time retrieval is disciplined by the objective itself:
183
+ # L_traj pins each intermediate state to its teacher segment, so a
184
+ # wait-then-fetch trajectory misses its intermediate targets. The gate
185
+ # bias initialises negative, i.e. training starts from the proven
186
+ # no-reread regime and departs only where gradients demand it.
187
+ # Default False: v4/v5 checkpoints reload bit-identically.
188
+ self.read_gating = read_gating
189
+ # Initial gate open-rate epsilon_gamma: b_gamma = logit(eps). Sigmoid
190
+ # never reaches exactly 0 at finite logits, so this is near-identity
191
+ # (not bitwise) initialisation w.r.t. the read-once path.
192
+ self.gate_init_eps = gate_init_eps
193
+ # v7: always-on PURE visual residual rereads. Every step adds
194
+ # Z = U + W_o XAttn(LN(U) -> queries, V* -> keys/values): no self-attn,
195
+ # no FFN, no output norm, no biases in the branch, W_o zero-init.
196
+ # Motivation (battery, ep1 ckpt 20752): read_once hurt badly while
197
+ # null_step (blank memory at one step) was harmless -- the gated
198
+ # reader's value was extra recurrent DEPTH, not vision. This branch
199
+ # removes every path that can change Z without V* content, so any
200
+ # benefit it shows IS visual by construction.
201
+ self.read_residual_pure = read_residual_pure
202
+ # v7.1 (advisor): relative-norm trust region on the visual residual.
203
+ # s_k = min(1, rho * RMS(U) / (RMS(Delta) + eps)); Z = U + s_k * Delta
204
+ # Unlike a fixed alpha, the optimizer cannot undo this by rescaling
205
+ # W_o (s_k adapts); ||Delta_used|| <~ rho * ||U|| structurally, and
206
+ # V* = 0 => Delta = 0 keeps the zero-memory identity. 0.0 disables.
207
+ self.residual_norm_cap = residual_norm_cap
208
+ # v7.1: initialise pure_attn K/V projections from the initial visual
209
+ # reader's v_in (same [d_w, d] shape) and Q from its first read
210
+ # block's in-proj slice, instead of random xavier -- the ep1 finding
211
+ # was that Q/K/V stayed ~random while only W_o trained.
212
+ self.pure_warm_start = pure_warm_start
213
+ # MAIN METHOD (advisor-confirmed): Final-Verified Recurrent Workspace.
214
+ # Initial read -> PURE latent recurrence (no mid-step visual anything)
215
+ # -> exactly ONE pure visual verification at the end:
216
+ # Z_out = U^(T) + W_o XAttn(LN(U^(T)), V*), alpha = 1, W_o = 0 init.
217
+ # Uses the same bias-free verifier module as read_residual_pure; the
218
+ # two flags differ only in WHERE the branch fires (every step vs once).
219
+ self.final_verify = final_verify
220
+ # MAIN (advisor 최종 확정): Mid-Read Recurrent Workspace. 고정 상수
221
+ # k* = floor(T/2) = 4 (스텝 sweep 산물이 아니라 판독 전후 동수 전이의
222
+ # midpoint)에서 pure visual read 1회:
223
+ # Z^(4) = U^(4) + W_o XAttn(LN(U^(4)), V*), alpha=1, W_o=0 init.
224
+ # 논리적 step 4 == python 루프 index 3 (0-based) -- off-by-one은
225
+ # tests/test_mid_read.py가 상태 분기 위치로 고정한다. 0이면 비활성.
226
+ self.mid_read_step = mid_read_step
227
+ # advisor(2026-08-06): E_slot(=workspace_read.slots)을 매 recurrent
228
+ # step 입력과 mid-read query에 일시 제공. state에는 누적하지 않는다.
229
+ self.persistent_slot_id = persistent_slot_id
230
+ # 고정 read 스케일 Z = U + alpha*W_o A (advisor scale-mismatch 검사).
231
+ # 1.0 = 기존과 비트동일. eval의 gate_overrides alpha는 절대값으로 이를 대체.
232
+ self.mid_read_alpha = mid_read_alpha
233
+ # advisor 최종안: E(persistent evidence) / R(recurrent reasoning) 분리.
234
+ # states = [E^(0), R^(1..T)], 결합 Z^(T)=E+R^(T)는 prefill에서.
235
+ self.evidence_reasoning = evidence_reasoning
236
+ # Current replacement candidate: no learned evidence/reasoning slots.
237
+ # E is the complete raw V* sequence and R0 is the text-only Q* sequence.
238
+ # A single shared CrossAttn transition produces R1..RT. Residual mode
239
+ # decodes RT; the replacement candidate decodes Q*+HT. Cross-attention
240
+ # uses the backbone's native GQA shape and
241
+ # is initialized from the first upper layer by default. Checkpoints
242
+ # created before the CrossAttn-only decision did not serialize
243
+ # raw_transition_ffn; None therefore preserves their historical FFN
244
+ # for faithful evaluation, while every new training launch passes False.
245
+ self.raw_evidence_question = bool(raw_evidence_question)
246
+ self.raw_attention_source_layer = int(raw_attention_source_layer)
247
+ self.raw_transition_ffn = (
248
+ True if raw_transition_ffn is None else bool(raw_transition_ffn)
249
+ )
250
+ # Experimental simplification of the same raw-E/Q path. The residual
251
+ # carrier is replaced, not augmented:
252
+ # H0=Q*, Hk=CrossAttn(H{k-1}, E), decoder_latent=Q*+HT.
253
+ # False preserves all residual raw-E/Q checkpoints and active runs.
254
+ self.raw_state_replacement = bool(raw_state_replacement)
255
+ # Question-anchored replacement keeps the fixed text scaffold in every
256
+ # recurrent query without restoring the old recurrent identity path:
257
+ # Z0=Q*, Zk=Q*+CrossAttn(Z{k-1}, E), decoder_latent=ZT.
258
+ # This replaces the one-time Q*+HT decoder interface above.
259
+ self.raw_question_anchor = bool(raw_question_anchor)
260
+ # Q-shaped recurrent candidate selected by the V* interface oracle:
261
+ # R0=Q*, Rk=R{k-1}+CrossAttn(R{k-1}, V*)
262
+ # R_agg=mean(R1..RT)
263
+ # This is parameter-free and replaces both question-anchored state
264
+ # replacement and final-state-only decoding.
265
+ self.raw_mean_aggregation = bool(raw_mean_aggregation)
266
+ # Counterfactual Residual Recurrence (CRR): the image-conditioned and
267
+ # text-only question states are propagated by the same native layer
268
+ # F=L(ell_star+1). Only their counterfactual residual is recurrently
269
+ # refined; there are no slots, V* memory, custom attention modules, or
270
+ # inference-time auxiliary modules. beta is the fixed relaxed-update
271
+ # coefficient; an optional batch-swap reliance loss is training-only.
272
+ self.counterfactual_residual_recurrence = bool(
273
+ counterfactual_residual_recurrence
274
+ )
275
+ # Qwen3 counterpart of the promoted Qwen2.5 full-state CRR path. The
276
+ # multimodal scaffold is used only inside the shared recurrent cell;
277
+ # answer decoding still receives a question-shaped state and a
278
+ # text-only prefix cache.
279
+ self.visual_counterfactual_recurrence = bool(
280
+ visual_counterfactual_recurrence
281
+ )
282
+ self.visual_cumulative_recurrence = bool(
283
+ visual_cumulative_recurrence
284
+ )
285
+ self.visual_full_state_recurrence = bool(
286
+ visual_full_state_recurrence
287
+ )
288
+ self.recurrence_only_adapter = bool(recurrence_only_adapter)
289
+ self.counterfactual_beta = float(counterfactual_beta)
290
+ # Bias-free Concat(C1..CT) -> C_agg, initialized to the exact
291
+ # historical mean interface for checkpoint-compatible warm starts.
292
+ self.crr_concat_aggregation = bool(crr_concat_aggregation)
293
+ # Optional latent-only policy adapter. It corrects the deterministic
294
+ # C2..CT means but never touches B, C1, or answer-token decoding. Rank 0
295
+ # is the exact historical CRR path; RL checkpoints serialize rank > 0.
296
+ self.crr_policy_rank = int(crr_policy_rank)
297
+ self.crr_policy_alpha = float(crr_policy_alpha)
298
+ # Training-only CRR intervention. A batch-matched residual from a
299
+ # different-answer example replaces the factual residual before the
300
+ # upper decoder. The model returns the factual-minus-swapped answer
301
+ # score gap; the Trainer applies the margin loss. Evaluation and
302
+ # generation never execute the extra branch.
303
+ self.counterfactual_reliance_loss = bool(
304
+ counterfactual_reliance_loss
305
+ )
306
+ self.transition_aware_reliance_loss = bool(
307
+ transition_aware_reliance_loss
308
+ )
309
+ # Training-only functional no-regression objective. Gold-answer NLL
310
+ # is measured through the actual upper decoder at every recurrent CRR
311
+ # prefix; inference remains the unchanged T-step mean interface.
312
+ self.counterfactual_step_retention_loss = bool(
313
+ counterfactual_step_retention_loss
314
+ )
315
+ # Training-only self-distillation target. The multimodal lower pass
316
+ # already computed for V* also provides its image-conditioned question
317
+ # rows. They supervise R_agg at the split boundary, but are never
318
+ # exposed to the inference decoder.
319
+ self.functional_state_loss = bool(functional_state_loss)
320
+ # Token CE lets long caption answers dominate one-token visual-choice
321
+ # examples. The per-example variant first averages valid answer tokens
322
+ # within each sample, then averages samples. It changes training/eval
323
+ # loss reduction only and has no inference-time effect.
324
+ self.per_example_answer_loss = bool(per_example_answer_loss)
325
+ # ER v2 (advisor 2026-08-11): explicit evidence read —
326
+ # R^(k) = f_theta(R^(k-1), Q*; E), E는 별도 memory로 cross-attn 읽기.
327
+ # (E+R 합산-감산 제거. evidence_reasoning=True 전제.)
328
+ self.explicit_evidence_read = bool(kwargs.pop("explicit_evidence_read", False))
329
+ # Historical field name retained for checkpoint compatibility. In ER
330
+ # v2.2 this means R^(0) is a batch-broadcast set of learnable slots; it
331
+ # is deliberately NOT initialized from Q*. Q* conditions every shared
332
+ # transition instead.
333
+ self.r_init_question = bool(kwargs.pop("r_init_question", False))
334
+ # ER v2.1 (advisor 2026-08-11): evidence read의 competitive normalization —
335
+ # slot 축 softmax 후 evidence 축 재정규화 (Slot Attention식 경쟁).
336
+ # 파라미터/모듈/loss 무추가; evidence_read의 W_qkv/W_o 재사용.
337
+ self.competitive_evidence_read = bool(
338
+ kwargs.pop("competitive_evidence_read", False))
339
+ # ER v2.2 C: k>=1 traj 감독 부재 시 twin 체인 제거 (states = answer 체인).
340
+ self.er_single_chain = bool(kwargs.pop("er_single_chain", False))
341
+ # ER v2.2 C (§5-7): decoder latent = softmax(temporal_logits)로 가중한
342
+ # R1..RT 결합 (전역 스칼라 T개). er_single_chain 전제. (legacy)
343
+ self.temporal_aggregation = bool(kwargs.pop("temporal_aggregation", False))
344
+ # ER v2.2 C 확정판: R_agg = Linear(Concat(R1..RT), bias=False),
345
+ # 초기 W=[I/T,...,I/T] (R_agg=mean(R1..RT)). er_single_chain 전제.
346
+ self.concat_aggregation = bool(kwargs.pop("concat_aggregation", False))
347
+ # ER v2: 워크스페이스 블록의 rank-r 부분공간 attention/FFN (0 = dense).
348
+ # d_w=d에서 state width 유지와 파라미터 폭증을 분리 (advisor).
349
+ self.workspace_low_rank = int(kwargs.pop("workspace_low_rank", 0))
350
+ # advisor: single-layer text-side interface alignment (bridge-only).
351
+ # teacher = frozen-base mm 경로 q_end@l*+1 (adapter off, sg);
352
+ # student = detach(R^T) splice의 l*+1 한 층 재계산. 추론 시 무존재.
353
+ self.interface_loss = interface_loss
354
+ # Training-only visual anchor switch. False in the current answer-only
355
+ # method; when false the model does not even materialize pooled V* for
356
+ # the Trainer.
357
+ self.read_anchor = bool(read_anchor)
358
+ if evidence_reasoning and persistent_slot_id:
359
+ raise ValueError("evidence_reasoning은 persistent_slot_id와 동시 사용 불가")
360
+ if evidence_reasoning and not mid_read_step:
361
+ raise ValueError("evidence_reasoning은 mid_read_step > 0 필요")
362
+ self.validate()
363
+
364
+ # -- derived views -----------------------------------------------------
365
+
366
+ @property
367
+ def num_decoder_layers(self) -> int:
368
+ return self.text_config.num_hidden_layers
369
+
370
+ @property
371
+ def backbone_width(self) -> int:
372
+ """``d`` -- backbone hidden width (4096 for the 7B)."""
373
+ return self.text_config.hidden_size
374
+
375
+ @property
376
+ def lower_slice(self) -> slice:
377
+ """Layers forming ``F_{<=l*}``."""
378
+ self._require_ell_star()
379
+ return slice(0, self.ell_star + 1)
380
+
381
+ @property
382
+ def upper_slice(self) -> slice:
383
+ """Layers forming ``F_{>l*}``."""
384
+ self._require_ell_star()
385
+ return slice(self.ell_star + 1, self.num_decoder_layers)
386
+
387
+ # -- validation --------------------------------------------------------
388
+
389
+ def _require_ell_star(self) -> None:
390
+ if self.ell_star is None:
391
+ raise ValueError(
392
+ "ell_star is unset. Run scripts/localize_read_point.py (Method 3.1) "
393
+ "and pass the measured layer explicitly; it must not be guessed."
394
+ )
395
+
396
+ def validate(self) -> None:
397
+ """Reject configurations that silently break the §3.2 contract."""
398
+ if self.ell_star is not None:
399
+ n = self.num_decoder_layers
400
+ # Both branches must be non-empty: an empty lower branch means there
401
+ # is no V* to read, an empty upper branch means the workspace never
402
+ # reaches the decoder.
403
+ if not 0 <= self.ell_star <= n - 2:
404
+ raise ValueError(
405
+ f"ell_star={self.ell_star} out of range for {n} decoder layers; "
406
+ f"expected 0 <= ell_star <= {n - 2} so both branches are non-empty."
407
+ )
408
+ # K=0 is the no-recurrence control (read only), not a misconfiguration.
409
+ if self.num_workspace_steps < 0:
410
+ raise ValueError("num_workspace_steps (K) must be >= 0.")
411
+ if self.num_workspace_slots < 1:
412
+ raise ValueError("num_workspace_slots (S) must be >= 1.")
413
+ if self.workspace_num_heads < 1:
414
+ raise ValueError("workspace_num_heads must be >= 1.")
415
+ if self.workspace_width % self.workspace_num_heads != 0:
416
+ raise ValueError(
417
+ f"workspace_width={self.workspace_width} must be divisible by "
418
+ f"workspace_num_heads={self.workspace_num_heads}."
419
+ )
420
+ if self.workspace_width > self.backbone_width:
421
+ raise ValueError(
422
+ f"workspace_width (d_w={self.workspace_width}) cannot exceed "
423
+ f"backbone width (d={self.backbone_width}). Native-width d_w=d is "
424
+ "the current method and uses an Identity decoder interface."
425
+ )
426
+ if self.workspace_rope_mode not in WORKSPACE_ROPE_MODES:
427
+ raise ValueError(
428
+ f"workspace_rope_mode={self.workspace_rope_mode!r} not in "
429
+ f"{WORKSPACE_ROPE_MODES}."
430
+ )
431
+ if self.adapter_rank < 0:
432
+ raise ValueError("adapter_rank must be >= 0 (0 disables the adapter).")
433
+ if not math.isfinite(self.adapter_dropout) or not 0.0 <= self.adapter_dropout < 1.0:
434
+ raise ValueError("adapter_dropout must be finite and in [0, 1).")
435
+ if getattr(self, "crr_policy_rank", 0) < 0:
436
+ raise ValueError("crr_policy_rank must be >= 0.")
437
+ if not math.isfinite(getattr(self, "crr_policy_alpha", 0.0)) or getattr(
438
+ self, "crr_policy_alpha", 0.0
439
+ ) <= 0.0:
440
+ raise ValueError("crr_policy_alpha must be finite and > 0.")
441
+ if self.read_num_blocks < 1 or self.transition_num_blocks < 1:
442
+ raise ValueError("read_num_blocks and transition_num_blocks must be >= 1.")
443
+ if self.workspace_ffn_mult < 1:
444
+ raise ValueError("workspace_ffn_mult must be >= 1.")
445
+ mid = int(getattr(self, "mid_read_step", 0) or 0)
446
+ if mid < 0 or mid > self.num_workspace_steps:
447
+ raise ValueError(
448
+ f"mid_read_step={mid} must be in [0, num_workspace_steps="
449
+ f"{self.num_workspace_steps}]."
450
+ )
451
+ low_rank = int(getattr(self, "workspace_low_rank", 0) or 0)
452
+ if low_rank < 0:
453
+ raise ValueError("workspace_low_rank must be >= 0.")
454
+ if low_rank and low_rank % self.workspace_num_heads != 0:
455
+ raise ValueError(
456
+ f"workspace_low_rank={low_rank} must be divisible by "
457
+ f"workspace_num_heads={self.workspace_num_heads}."
458
+ )
459
+ if getattr(self, "competitive_evidence_read", False) and not getattr(
460
+ self, "explicit_evidence_read", False
461
+ ):
462
+ raise ValueError(
463
+ "competitive_evidence_read requires explicit_evidence_read=True; "
464
+ "otherwise the flag has no execution path."
465
+ )
466
+ if getattr(self, "explicit_evidence_read", False) and not getattr(
467
+ self, "evidence_reasoning", False
468
+ ):
469
+ raise ValueError(
470
+ "explicit_evidence_read requires evidence_reasoning=True; otherwise "
471
+ "the evidence-read module is allocated but never called."
472
+ )
473
+ if getattr(self, "r_init_question", False) and not getattr(
474
+ self, "evidence_reasoning", False
475
+ ):
476
+ raise ValueError(
477
+ "r_init_question (legacy name for learnable R0 slots) requires "
478
+ "evidence_reasoning=True; otherwise reasoning_slots are unused."
479
+ )
480
+ if (getattr(self, "concat_aggregation", False)
481
+ or getattr(self, "temporal_aggregation", False)):
482
+ if self.num_workspace_steps < 1:
483
+ raise ValueError("reasoning aggregation requires num_workspace_steps >= 1.")
484
+ if not getattr(self, "evidence_reasoning", False):
485
+ raise ValueError(
486
+ "reasoning aggregation requires evidence_reasoning=True; "
487
+ "otherwise the aggregation module is unused."
488
+ )
489
+ if getattr(self, "concat_aggregation", False) and not getattr(
490
+ self, "er_single_chain", False
491
+ ):
492
+ raise ValueError("concat_aggregation requires er_single_chain=True.")
493
+ if getattr(self, "raw_evidence_question", False):
494
+ if self.workspace_width != self.backbone_width:
495
+ raise ValueError(
496
+ "raw_evidence_question requires native workspace_width == "
497
+ "backbone_width: E=V* and R0=Q* are used without projections."
498
+ )
499
+ if self.num_workspace_steps < 1:
500
+ raise ValueError(
501
+ "raw_evidence_question requires num_workspace_steps >= 1."
502
+ )
503
+ source_layer = (
504
+ self.ell_star + 1
505
+ if self.raw_attention_source_layer < 0
506
+ else self.raw_attention_source_layer
507
+ )
508
+ if not self.ell_star < source_layer < self.text_config.num_hidden_layers:
509
+ raise ValueError(
510
+ "raw_attention_source_layer must be an upper-backbone layer, "
511
+ f"got ell_star={self.ell_star}, source={source_layer}."
512
+ )
513
+ incompatible = {
514
+ "evidence_reasoning": bool(getattr(self, "evidence_reasoning", False)),
515
+ "explicit_evidence_read": bool(getattr(self, "explicit_evidence_read", False)),
516
+ "competitive_evidence_read": bool(getattr(self, "competitive_evidence_read", False)),
517
+ "r_init_question": bool(getattr(self, "r_init_question", False)),
518
+ "er_single_chain": bool(getattr(self, "er_single_chain", False)),
519
+ "temporal_aggregation": bool(getattr(self, "temporal_aggregation", False)),
520
+ "concat_aggregation": bool(getattr(self, "concat_aggregation", False)),
521
+ "read_shortcut": bool(getattr(self, "read_shortcut", False)),
522
+ "splice_scale_match": bool(getattr(self, "splice_scale_match", False)),
523
+ "read_gating": bool(getattr(self, "read_gating", False)),
524
+ "read_residual_pure": bool(getattr(self, "read_residual_pure", False)),
525
+ "final_verify": bool(getattr(self, "final_verify", False)),
526
+ "mid_read_step": bool(getattr(self, "mid_read_step", 0)),
527
+ "persistent_slot_id": bool(getattr(self, "persistent_slot_id", False)),
528
+ "interface_loss": bool(getattr(self, "interface_loss", False)),
529
+ "read_anchor": bool(getattr(self, "read_anchor", False)),
530
+ }
531
+ active = [name for name, enabled in incompatible.items() if enabled]
532
+ if active:
533
+ raise ValueError(
534
+ "raw_evidence_question replaces the slot workspace and is "
535
+ f"incompatible with: {', '.join(active)}"
536
+ )
537
+ if self.raw_state_replacement and self.raw_transition_ffn:
538
+ raise ValueError(
539
+ "raw_state_replacement is CrossAttn-only and is incompatible "
540
+ "with raw_transition_ffn=True."
541
+ )
542
+ if self.raw_question_anchor and not self.raw_state_replacement:
543
+ raise ValueError(
544
+ "raw_question_anchor requires raw_state_replacement=True."
545
+ )
546
+ if self.raw_mean_aggregation and (
547
+ self.raw_state_replacement or self.raw_question_anchor
548
+ ):
549
+ raise ValueError(
550
+ "raw_mean_aggregation replaces raw_state_replacement and "
551
+ "raw_question_anchor; use the residual recurrence."
552
+ )
553
+ if self.functional_state_loss and not self.raw_mean_aggregation:
554
+ raise ValueError(
555
+ "functional_state_loss requires raw_mean_aggregation=True "
556
+ "so its target is the exact decoder-facing recurrence."
557
+ )
558
+ elif getattr(self, "raw_state_replacement", False):
559
+ raise ValueError(
560
+ "raw_state_replacement requires raw_evidence_question=True."
561
+ )
562
+ elif getattr(self, "raw_question_anchor", False):
563
+ raise ValueError(
564
+ "raw_question_anchor requires raw_evidence_question=True."
565
+ )
566
+ elif getattr(self, "raw_mean_aggregation", False):
567
+ raise ValueError(
568
+ "raw_mean_aggregation requires raw_evidence_question=True."
569
+ )
570
+ elif getattr(self, "functional_state_loss", False):
571
+ raise ValueError(
572
+ "functional_state_loss requires raw_evidence_question=True."
573
+ )
574
+ if (
575
+ getattr(self, "counterfactual_reliance_loss", False)
576
+ and not getattr(self, "counterfactual_residual_recurrence", False)
577
+ ):
578
+ raise ValueError(
579
+ "counterfactual_reliance_loss requires "
580
+ "counterfactual_residual_recurrence=True."
581
+ )
582
+ if (
583
+ getattr(self, "transition_aware_reliance_loss", False)
584
+ and not getattr(self, "counterfactual_reliance_loss", False)
585
+ ):
586
+ raise ValueError(
587
+ "transition_aware_reliance_loss requires "
588
+ "counterfactual_reliance_loss=True."
589
+ )
590
+ if (
591
+ getattr(self, "transition_aware_reliance_loss", False)
592
+ and self.num_workspace_steps < 2
593
+ ):
594
+ raise ValueError(
595
+ "transition_aware_reliance_loss requires at least two steps."
596
+ )
597
+ if (
598
+ getattr(self, "counterfactual_step_retention_loss", False)
599
+ and not getattr(self, "counterfactual_residual_recurrence", False)
600
+ ):
601
+ raise ValueError(
602
+ "counterfactual_step_retention_loss requires "
603
+ "counterfactual_residual_recurrence=True."
604
+ )
605
+ if (
606
+ getattr(self, "counterfactual_step_retention_loss", False)
607
+ and self.num_workspace_steps < 2
608
+ ):
609
+ raise ValueError(
610
+ "counterfactual_step_retention_loss requires at least two steps."
611
+ )
612
+ if getattr(self, "crr_concat_aggregation", False) and not getattr(
613
+ self, "counterfactual_residual_recurrence", False
614
+ ):
615
+ raise ValueError(
616
+ "crr_concat_aggregation requires "
617
+ "counterfactual_residual_recurrence=True."
618
+ )
619
+ if getattr(self, "counterfactual_residual_recurrence", False):
620
+ if getattr(self, "raw_evidence_question", False):
621
+ raise ValueError(
622
+ "counterfactual_residual_recurrence replaces "
623
+ "raw_evidence_question; enable exactly one q-state method."
624
+ )
625
+ if self.num_workspace_steps < 1:
626
+ raise ValueError(
627
+ "counterfactual_residual_recurrence requires "
628
+ "num_workspace_steps >= 1."
629
+ )
630
+ if self.crr_policy_rank > 0 and self.num_workspace_steps < 2:
631
+ raise ValueError(
632
+ "crr_policy_rank > 0 requires num_workspace_steps >= 2."
633
+ )
634
+ if not 0.0 <= self.counterfactual_beta <= 1.0:
635
+ raise ValueError(
636
+ "counterfactual_beta must lie in [0, 1], got "
637
+ f"{self.counterfactual_beta}."
638
+ )
639
+ incompatible = {
640
+ "evidence_reasoning": bool(getattr(self, "evidence_reasoning", False)),
641
+ "explicit_evidence_read": bool(getattr(self, "explicit_evidence_read", False)),
642
+ "competitive_evidence_read": bool(getattr(self, "competitive_evidence_read", False)),
643
+ "r_init_question": bool(getattr(self, "r_init_question", False)),
644
+ "er_single_chain": bool(getattr(self, "er_single_chain", False)),
645
+ "temporal_aggregation": bool(getattr(self, "temporal_aggregation", False)),
646
+ "concat_aggregation": bool(getattr(self, "concat_aggregation", False)),
647
+ "read_shortcut": bool(getattr(self, "read_shortcut", False)),
648
+ "splice_scale_match": bool(getattr(self, "splice_scale_match", False)),
649
+ "read_gating": bool(getattr(self, "read_gating", False)),
650
+ "read_residual_pure": bool(getattr(self, "read_residual_pure", False)),
651
+ "final_verify": bool(getattr(self, "final_verify", False)),
652
+ "mid_read_step": bool(getattr(self, "mid_read_step", 0)),
653
+ "persistent_slot_id": bool(getattr(self, "persistent_slot_id", False)),
654
+ "interface_loss": bool(getattr(self, "interface_loss", False)),
655
+ "read_anchor": bool(getattr(self, "read_anchor", False)),
656
+ "functional_state_loss": bool(getattr(self, "functional_state_loss", False)),
657
+ }
658
+ active = [name for name, enabled in incompatible.items() if enabled]
659
+ if active:
660
+ raise ValueError(
661
+ "counterfactual_residual_recurrence replaces the slot/raw "
662
+ f"workspace and is incompatible with: {', '.join(active)}"
663
+ )
664
+ elif getattr(self, "crr_policy_rank", 0) > 0:
665
+ raise ValueError(
666
+ "crr_policy_rank > 0 requires "
667
+ "counterfactual_residual_recurrence=True."
668
+ )
669
+ if getattr(self, "visual_counterfactual_recurrence", False):
670
+ if not getattr(self, "counterfactual_residual_recurrence", False):
671
+ raise ValueError(
672
+ "visual_counterfactual_recurrence requires "
673
+ "counterfactual_residual_recurrence=True."
674
+ )
675
+ if not getattr(self, "visual_cumulative_recurrence", False):
676
+ raise ValueError(
677
+ "Qwen3 visual recurrence currently requires the promoted "
678
+ "cumulative full-state path."
679
+ )
680
+ if not getattr(self, "visual_full_state_recurrence", False):
681
+ raise ValueError(
682
+ "Qwen3 visual recurrence currently requires "
683
+ "visual_full_state_recurrence=True."
684
+ )
685
+ incompatible = {
686
+ "counterfactual_reliance_loss": bool(
687
+ getattr(self, "counterfactual_reliance_loss", False)
688
+ ),
689
+ "transition_aware_reliance_loss": bool(
690
+ getattr(self, "transition_aware_reliance_loss", False)
691
+ ),
692
+ "counterfactual_step_retention_loss": bool(
693
+ getattr(self, "counterfactual_step_retention_loss", False)
694
+ ),
695
+ "crr_concat_aggregation": bool(
696
+ getattr(self, "crr_concat_aggregation", False)
697
+ ),
698
+ "crr_policy_rank": bool(getattr(self, "crr_policy_rank", 0)),
699
+ }
700
+ active = [name for name, enabled in incompatible.items() if enabled]
701
+ if active:
702
+ raise ValueError(
703
+ "Qwen3 full-state visual recurrence uses final-answer CE "
704
+ "and final-state decoding only; incompatible with: "
705
+ + ", ".join(active)
706
+ )
707
+ elif getattr(self, "visual_cumulative_recurrence", False) or getattr(
708
+ self, "visual_full_state_recurrence", False
709
+ ):
710
+ raise ValueError(
711
+ "visual cumulative/full-state flags require "
712
+ "visual_counterfactual_recurrence=True."
713
+ )
714
+ if getattr(self, "recurrence_only_adapter", False):
715
+ if not getattr(self, "visual_full_state_recurrence", False):
716
+ raise ValueError(
717
+ "recurrence_only_adapter requires "
718
+ "visual_full_state_recurrence=True."
719
+ )
720
+ if self.adapter_rank <= 0:
721
+ raise ValueError(
722
+ "recurrence_only_adapter requires adapter_rank > 0."
723
+ )
724
+ recurrent_layer = int(self.ell_star) + 1
725
+ if recurrent_layer in set(self.adapter_exclude_layers):
726
+ raise ValueError(
727
+ "recurrence_only_adapter cannot exclude its recurrent "
728
+ f"layer {recurrent_layer} from LoRA."
729
+ )
730
+ modes = [bool(getattr(self, m, False)) for m in
731
+ ("read_gating", "read_residual_pure", "final_verify")]
732
+ modes.append(bool(getattr(self, "mid_read_step", 0)))
733
+ if sum(modes) > 1:
734
+ raise ValueError(
735
+ "read_gating / read_residual_pure / final_verify are mutually "
736
+ "exclusive reread modes."
737
+ )
738
+
739
+
740
+ __all__ = ["CloseQwen3VLConfig", "WORKSPACE_ROPE_MODES"]
cvrr_release_config.json ADDED
@@ -0,0 +1,19 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "name": "CVRR-InternVL3-9B",
3
+ "base_model": "OpenGVLab/InternVL3-9B",
4
+ "ell_star": 35,
5
+ "recurrent_layer": 36,
6
+ "upper_decoder_start": 37,
7
+ "inference_T": 4,
8
+ "inference_beta": 0.33,
9
+ "lora_rank": 32,
10
+ "lora_alpha": 12.0,
11
+ "lora_dropout": 0.01,
12
+ "checkpoint_step": 500,
13
+ "format": "cvrr_native_plus_merged_transition_v1",
14
+ "merged_weight_dtype": "float32",
15
+ "native_weight_bytes": 18277586944,
16
+ "status": "gpu_vstar_comparison_completed",
17
+ "upload_ready": true,
18
+ "hub_model_id": "dmis-lab/InternVL3-9B-CVRR"
19
+ }
licenses/Apache-2.0.txt ADDED
@@ -0,0 +1,202 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+
2
+ Apache License
3
+ Version 2.0, January 2004
4
+ http://www.apache.org/licenses/
5
+
6
+ TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
7
+
8
+ 1. Definitions.
9
+
10
+ "License" shall mean the terms and conditions for use, reproduction,
11
+ and distribution as defined by Sections 1 through 9 of this document.
12
+
13
+ "Licensor" shall mean the copyright owner or entity authorized by
14
+ the copyright owner that is granting the License.
15
+
16
+ "Legal Entity" shall mean the union of the acting entity and all
17
+ other entities that control, are controlled by, or are under common
18
+ control with that entity. For the purposes of this definition,
19
+ "control" means (i) the power, direct or indirect, to cause the
20
+ direction or management of such entity, whether by contract or
21
+ otherwise, or (ii) ownership of fifty percent (50%) or more of the
22
+ outstanding shares, or (iii) beneficial ownership of such entity.
23
+
24
+ "You" (or "Your") shall mean an individual or Legal Entity
25
+ exercising permissions granted by this License.
26
+
27
+ "Source" form shall mean the preferred form for making modifications,
28
+ including but not limited to software source code, documentation
29
+ source, and configuration files.
30
+
31
+ "Object" form shall mean any form resulting from mechanical
32
+ transformation or translation of a Source form, including but
33
+ not limited to compiled object code, generated documentation,
34
+ and conversions to other media types.
35
+
36
+ "Work" shall mean the work of authorship, whether in Source or
37
+ Object form, made available under the License, as indicated by a
38
+ copyright notice that is included in or attached to the work
39
+ (an example is provided in the Appendix below).
40
+
41
+ "Derivative Works" shall mean any work, whether in Source or Object
42
+ form, that is based on (or derived from) the Work and for which the
43
+ editorial revisions, annotations, elaborations, or other modifications
44
+ represent, as a whole, an original work of authorship. For the purposes
45
+ of this License, Derivative Works shall not include works that remain
46
+ separable from, or merely link (or bind by name) to the interfaces of,
47
+ the Work and Derivative Works thereof.
48
+
49
+ "Contribution" shall mean any work of authorship, including
50
+ the original version of the Work and any modifications or additions
51
+ to that Work or Derivative Works thereof, that is intentionally
52
+ submitted to Licensor for inclusion in the Work by the copyright owner
53
+ or by an individual or Legal Entity authorized to submit on behalf of
54
+ the copyright owner. For the purposes of this definition, "submitted"
55
+ means any form of electronic, verbal, or written communication sent
56
+ to the Licensor or its representatives, including but not limited to
57
+ communication on electronic mailing lists, source code control systems,
58
+ and issue tracking systems that are managed by, or on behalf of, the
59
+ Licensor for the purpose of discussing and improving the Work, but
60
+ excluding communication that is conspicuously marked or otherwise
61
+ designated in writing by the copyright owner as "Not a Contribution."
62
+
63
+ "Contributor" shall mean Licensor and any individual or Legal Entity
64
+ on behalf of whom a Contribution has been received by Licensor and
65
+ subsequently incorporated within the Work.
66
+
67
+ 2. Grant of Copyright License. Subject to the terms and conditions of
68
+ this License, each Contributor hereby grants to You a perpetual,
69
+ worldwide, non-exclusive, no-charge, royalty-free, irrevocable
70
+ copyright license to reproduce, prepare Derivative Works of,
71
+ publicly display, publicly perform, sublicense, and distribute the
72
+ Work and such Derivative Works in Source or Object form.
73
+
74
+ 3. Grant of Patent License. Subject to the terms and conditions of
75
+ this License, each Contributor hereby grants to You a perpetual,
76
+ worldwide, non-exclusive, no-charge, royalty-free, irrevocable
77
+ (except as stated in this section) patent license to make, have made,
78
+ use, offer to sell, sell, import, and otherwise transfer the Work,
79
+ where such license applies only to those patent claims licensable
80
+ by such Contributor that are necessarily infringed by their
81
+ Contribution(s) alone or by combination of their Contribution(s)
82
+ with the Work to which such Contribution(s) was submitted. If You
83
+ institute patent litigation against any entity (including a
84
+ cross-claim or counterclaim in a lawsuit) alleging that the Work
85
+ or a Contribution incorporated within the Work constitutes direct
86
+ or contributory patent infringement, then any patent licenses
87
+ granted to You under this License for that Work shall terminate
88
+ as of the date such litigation is filed.
89
+
90
+ 4. Redistribution. You may reproduce and distribute copies of the
91
+ Work or Derivative Works thereof in any medium, with or without
92
+ modifications, and in Source or Object form, provided that You
93
+ meet the following conditions:
94
+
95
+ (a) You must give any other recipients of the Work or
96
+ Derivative Works a copy of this License; and
97
+
98
+ (b) You must cause any modified files to carry prominent notices
99
+ stating that You changed the files; and
100
+
101
+ (c) You must retain, in the Source form of any Derivative Works
102
+ that You distribute, all copyright, patent, trademark, and
103
+ attribution notices from the Source form of the Work,
104
+ excluding those notices that do not pertain to any part of
105
+ the Derivative Works; and
106
+
107
+ (d) If the Work includes a "NOTICE" text file as part of its
108
+ distribution, then any Derivative Works that You distribute must
109
+ include a readable copy of the attribution notices contained
110
+ within such NOTICE file, excluding those notices that do not
111
+ pertain to any part of the Derivative Works, in at least one
112
+ of the following places: within a NOTICE text file distributed
113
+ as part of the Derivative Works; within the Source form or
114
+ documentation, if provided along with the Derivative Works; or,
115
+ within a display generated by the Derivative Works, if and
116
+ wherever such third-party notices normally appear. The contents
117
+ of the NOTICE file are for informational purposes only and
118
+ do not modify the License. You may add Your own attribution
119
+ notices within Derivative Works that You distribute, alongside
120
+ or as an addendum to the NOTICE text from the Work, provided
121
+ that such additional attribution notices cannot be construed
122
+ as modifying the License.
123
+
124
+ You may add Your own copyright statement to Your modifications and
125
+ may provide additional or different license terms and conditions
126
+ for use, reproduction, or distribution of Your modifications, or
127
+ for any such Derivative Works as a whole, provided Your use,
128
+ reproduction, and distribution of the Work otherwise complies with
129
+ the conditions stated in this License.
130
+
131
+ 5. Submission of Contributions. Unless You explicitly state otherwise,
132
+ any Contribution intentionally submitted for inclusion in the Work
133
+ by You to the Licensor shall be under the terms and conditions of
134
+ this License, without any additional terms or conditions.
135
+ Notwithstanding the above, nothing herein shall supersede or modify
136
+ the terms of any separate license agreement you may have executed
137
+ with Licensor regarding such Contributions.
138
+
139
+ 6. Trademarks. This License does not grant permission to use the trade
140
+ names, trademarks, service marks, or product names of the Licensor,
141
+ except as required for reasonable and customary use in describing the
142
+ origin of the Work and reproducing the content of the NOTICE file.
143
+
144
+ 7. Disclaimer of Warranty. Unless required by applicable law or
145
+ agreed to in writing, Licensor provides the Work (and each
146
+ Contributor provides its Contributions) on an "AS IS" BASIS,
147
+ WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
148
+ implied, including, without limitation, any warranties or conditions
149
+ of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
150
+ PARTICULAR PURPOSE. You are solely responsible for determining the
151
+ appropriateness of using or redistributing the Work and assume any
152
+ risks associated with Your exercise of permissions under this License.
153
+
154
+ 8. Limitation of Liability. In no event and under no legal theory,
155
+ whether in tort (including negligence), contract, or otherwise,
156
+ unless required by applicable law (such as deliberate and grossly
157
+ negligent acts) or agreed to in writing, shall any Contributor be
158
+ liable to You for damages, including any direct, indirect, special,
159
+ incidental, or consequential damages of any character arising as a
160
+ result of this License or out of the use or inability to use the
161
+ Work (including but not limited to damages for loss of goodwill,
162
+ work stoppage, computer failure or malfunction, or any and all
163
+ other commercial damages or losses), even if such Contributor
164
+ has been advised of the possibility of such damages.
165
+
166
+ 9. Accepting Warranty or Additional Liability. While redistributing
167
+ the Work or Derivative Works thereof, You may choose to offer,
168
+ and charge a fee for, acceptance of support, warranty, indemnity,
169
+ or other liability obligations and/or rights consistent with this
170
+ License. However, in accepting such obligations, You may act only
171
+ on Your own behalf and on Your sole responsibility, not on behalf
172
+ of any other Contributor, and only if You agree to indemnify,
173
+ defend, and hold each Contributor harmless for any liability
174
+ incurred by, or claims asserted against, such Contributor by reason
175
+ of your accepting any such warranty or additional liability.
176
+
177
+ END OF TERMS AND CONDITIONS
178
+
179
+ APPENDIX: How to apply the Apache License to your work.
180
+
181
+ To apply the Apache License to your work, attach the following
182
+ boilerplate notice, with the fields enclosed by brackets "[]"
183
+ replaced with your own identifying information. (Don't include
184
+ the brackets!) The text should be enclosed in the appropriate
185
+ comment syntax for the file format. We also recommend that a
186
+ file or class name and description of purpose be included on the
187
+ same "printed page" as the copyright notice for easier
188
+ identification within third-party archives.
189
+
190
+ Copyright [yyyy] [name of copyright owner]
191
+
192
+ Licensed under the Apache License, Version 2.0 (the "License");
193
+ you may not use this file except in compliance with the License.
194
+ You may obtain a copy of the License at
195
+
196
+ http://www.apache.org/licenses/LICENSE-2.0
197
+
198
+ Unless required by applicable law or agreed to in writing, software
199
+ distributed under the License is distributed on an "AS IS" BASIS,
200
+ WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
201
+ See the License for the specific language governing permissions and
202
+ limitations under the License.
licenses/InternVL-MIT.txt ADDED
@@ -0,0 +1,21 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ MIT License
2
+
3
+ Copyright (c) 2023 OpenGVLab
4
+
5
+ Permission is hereby granted, free of charge, to any person obtaining a copy
6
+ of this software and associated documentation files (the "Software"), to deal
7
+ in the Software without restriction, including without limitation the rights
8
+ to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
9
+ copies of the Software, and to permit persons to whom the Software is
10
+ furnished to do so, subject to the following conditions:
11
+
12
+ The above copyright notice and this permission notice shall be included in all
13
+ copies or substantial portions of the Software.
14
+
15
+ THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
16
+ IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
17
+ FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
18
+ AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
19
+ LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
20
+ OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
21
+ SOFTWARE.
licenses/UPSTREAM_MODEL_CARD.md ADDED
@@ -0,0 +1,701 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ license: mit
3
+ pipeline_tag: image-text-to-text
4
+ library_name: transformers
5
+ base_model:
6
+ - OpenGVLab/InternVL3-9B-Instruct
7
+ base_model_relation: finetune
8
+ datasets:
9
+ - OpenGVLab/MMPR-v1.2
10
+ language:
11
+ - multilingual
12
+ tags:
13
+ - internvl
14
+ - custom_code
15
+ ---
16
+
17
+ # InternVL3-9B
18
+
19
+ [\[📂 GitHub\]](https://github.com/OpenGVLab/InternVL) [\[📜 InternVL 1.0\]](https://huggingface.co/papers/2312.14238) [\[📜 InternVL 1.5\]](https://huggingface.co/papers/2404.16821) [\[📜 InternVL 2.5\]](https://huggingface.co/papers/2412.05271) [\[📜 InternVL2.5-MPO\]](https://huggingface.co/papers/2411.10442) [\[📜 InternVL3\]](https://huggingface.co/papers/2504.10479)
20
+
21
+ [\[🆕 Blog\]](https://internvl.github.io/blog/) [\[🗨️ Chat Demo\]](https://internvl.opengvlab.com/) [\[🤗 HF Demo\]](https://huggingface.co/spaces/OpenGVLab/InternVL) [\[🚀 Quick Start\]](#quick-start) [\[📖 Documents\]](https://internvl.readthedocs.io/en/latest/)
22
+
23
+ <div align="center">
24
+ <img width="500" alt="image" src="https://cdn-uploads.huggingface.co/production/uploads/64006c09330a45b03605bba3/zJsd2hqd3EevgXo6fNgC-.png">
25
+ </div>
26
+
27
+ ## Introduction
28
+
29
+ We introduce InternVL3, an advanced multimodal large language model (MLLM) series that demonstrates superior overall performance.
30
+ Compared to InternVL 2.5, InternVL3 exhibits superior multimodal perception and reasoning capabilities, while further extending its multimodal capabilities to encompass tool usage, GUI agents, industrial image analysis, 3D vision perception, and more.
31
+ Additionally, we compare InternVL3 with Qwen2.5 Chat models, whose corresponding pre-trained base models are employed as the initialization of the langauge component in InternVL3. Benefitting from Native Multimodal Pre-Training, the InternVL3 series achieves even better overall text performance than the Qwen2.5 series.
32
+
33
+ ![image/png](https://huggingface.co/datasets/Weiyun1025/InternVL-Performance/resolve/main/internvl3/overall.png)
34
+
35
+ ## InternVL3 Family
36
+
37
+ In the following table, we provide an overview of the InternVL3 series.
38
+
39
+ | Model Name | Vision Part | Language Part | HF Link |
40
+ | :-----------: | :-------------------------------------------------------------------------------------: | :----------------------------------------------------------------------------: | :------------------------------------------------------: |
41
+ | InternVL3-1B | [InternViT-300M-448px-V2_5](https://huggingface.co/OpenGVLab/InternViT-300M-448px-V2_5) | [Qwen2.5-0.5B](https://huggingface.co/Qwen/Qwen2.5-0.5B) | [🤗 link](https://huggingface.co/OpenGVLab/InternVL3-1B) |
42
+ | InternVL3-2B | [InternViT-300M-448px-V2_5](https://huggingface.co/OpenGVLab/InternViT-300M-448px-V2_5) | [Qwen2.5-1.5B](https://huggingface.co/Qwen/Qwen2.5-1.5B) | [🤗 link](https://huggingface.co/OpenGVLab/InternVL3-2B) |
43
+ | InternVL3-8B | [InternViT-300M-448px-V2_5](https://huggingface.co/OpenGVLab/InternViT-300M-448px-V2_5) | [Qwen2.5-7B](https://huggingface.co/Qwen/Qwen2.5-7B) | [🤗 link](https://huggingface.co/OpenGVLab/InternVL3-8B) |
44
+ | InternVL3-9B | [InternViT-300M-448px-V2_5](https://huggingface.co/OpenGVLab/InternViT-300M-448px-V2_5) | [internlm3-8b-instruct](https://huggingface.co/internlm/internlm3-8b-instruct) | [🤗 link](https://huggingface.co/OpenGVLab/InternVL3-9B) |
45
+ | InternVL3-14B | [InternViT-300M-448px-V2_5](https://huggingface.co/OpenGVLab/InternViT-300M-448px-V2_5) | [Qwen2.5-14B](https://huggingface.co/Qwen/Qwen2.5-14B) | [🤗 link](https://huggingface.co/OpenGVLab/InternVL3-14B) |
46
+ | InternVL3-38B | [InternViT-6B-448px-V2_5](https://huggingface.co/OpenGVLab/InternViT-6B-448px-V2_5) | [Qwen2.5-32B](https://huggingface.co/Qwen/Qwen2.5-32B) | [🤗 link](https://huggingface.co/OpenGVLab/InternVL3-38B) |
47
+ | InternVL3-78B | [InternViT-6B-448px-V2_5](https://huggingface.co/OpenGVLab/InternViT-6B-448px-V2_5) | [Qwen2.5-72B](https://huggingface.co/Qwen/Qwen2.5-72B) | [🤗 link](https://huggingface.co/OpenGVLab/InternVL3-78B) |
48
+
49
+ ![image/png](https://huggingface.co/datasets/Weiyun1025/InternVL-Performance/resolve/main/internvl3/overall-table.png)
50
+
51
+ ## Model Architecture
52
+
53
+ As shown in the following figure, [InternVL3](https://internvl.github.io/blog/2025-04-11-InternVL-3/) retains the same model architecture as [InternVL 2.5](https://internvl.github.io/blog/2024-12-05-InternVL-2.5/) and its predecessors, InternVL 1.5 and 2.0, following the "ViT-MLP-LLM" paradigm. In this new version, we integrate a newly incrementally pre-trained InternViT with various pre-trained LLMs, including InternLM 3 and Qwen 2.5, using a randomly initialized MLP projector.
54
+
55
+
56
+ ![image/png](https://cdn-uploads.huggingface.co/production/uploads/64119264f0f81eb569e0d569/BiiyXN6NOk0p-3rl3ueyL.png)
57
+
58
+ As in the previous version, we applied a pixel unshuffle operation, reducing the number of visual tokens to one-quarter of the original. Besides, we adopted a similar dynamic resolution strategy as InternVL 1.5, dividing images into tiles of 448×448 pixels. The key difference, starting from InternVL 2.0, is that we additionally introduced support for multi-image and video data.
59
+
60
+ Notably, in InternVL3, we integrate the [Variable Visual Position Encoding (V2PE)](https://arxiv.org/abs/2412.09616), which utilizes smaller, more flexible position increments for visual tokens. Benefiting from V2PE, InternVL3 exhibits better long context understanding capabilities compared to its predecessors.
61
+
62
+ ## Training Strategy
63
+
64
+ ### Native Multimodal Pre-Training
65
+
66
+ We propose a [Native Multimodal Pre-Training](https://huggingface.co/papers/2504.10479) approach that consolidates language and vision learning into a single pre-training stage.
67
+ In contrast to standard paradigms that first train a language-only model and subsequently adapt it to handle additional modalities, our method interleaves multimodal data (e.g., image-text, video-text, or image-text interleaved sequences) with large-scale textual corpora. This unified training scheme allows the model to learn both linguistic and multimodal representations simultaneously, ultimately enhancing its capability to handle vision-language tasks without the need for separate alignment or bridging modules.
68
+ Please see [our paper](https://huggingface.co/papers/2504.10479) for more details.
69
+
70
+ ### Supervised Fine-Tuning
71
+
72
+ In this phase, the techniques of random JPEG compression, square loss re-weighting, and multimodal data packing proposed in [InternVL2.5](https://arxiv.org/abs/2412.05271) are also employed in the InternVL3 series.
73
+ The main advancement of the SFT phase in InternVL3 compared to InternVL2.5 lies in the use of higher-quality and more diverse training data.
74
+ Specifically, we further extend training samples for tool use, 3D scene understanding, GUI operations, long context tasks, video understanding, scientific diagrams, creative writing, and multimodal reasoning.
75
+
76
+ ### Mixed Preference Optimization
77
+
78
+ During Pre-training and SFT, the model is trained to predict the next token conditioned on previous ground-truth tokens.
79
+ However, during inference, the model predicts each token based on its own prior outputs.
80
+ This discrepancy between ground-truth tokens and model-predicted tokens introduces a distribution shift, which can impair the model’s Chain-of-Thought (CoT) reasoning capabilities.
81
+ To mitigate this issue, we employ [MPO](https://arxiv.org/abs/2411.10442), which introduces additional supervision from both positive and negative samples to align the model response distribution with the ground-truth distribution, thereby improving reasoning performance.
82
+ Specifically, the training objective of MPO is a combination of
83
+ preference loss \\(\mathcal{L}_{\text{p}}\\),
84
+ quality loss \\(\mathcal{L}_{\text{q}}\\),
85
+ and generation loss \\(\mathcal{L}_{\text{g}}\\),
86
+ which can be formulated as follows:
87
+
88
+
89
+ $$
90
+ \mathcal{L}=w_{p}\cdot\mathcal{L}_{\text{p}} + w_{q}\cdot\mathcal{L}_{\text{q}} + w_{g}\cdot\mathcal{L}_{\text{g}},
91
+ $$
92
+
93
+
94
+ where \\(w_{*}\\) represents the weight assigned to each loss component. Please see [our paper](https://arxiv.org/abs/2411.10442) for more details about MPO.
95
+
96
+
97
+ ### Test-Time Scaling
98
+
99
+ Test-Time Scaling has been shown to be an effective method to enhance the reasoning abilities of LLMs and MLLMs.
100
+ In this work, we use the Best-of-N evaluation strategy and employ [VisualPRM-8B](https://huggingface.co/OpenGVLab/VisualPRM-8B) as the critic model to select the best response for reasoning and mathematics evaluation.
101
+
102
+ ## Evaluation on Multimodal Capability
103
+
104
+ ### Multimodal Reasoning and Mathematics
105
+
106
+ ![image/png](https://huggingface.co/OpenGVLab/VisualPRM-8B-v1_1/resolve/main/visualprm-performance.png)
107
+
108
+ ### OCR, Chart, and Document Understanding
109
+
110
+ ![image/png](https://huggingface.co/datasets/Weiyun1025/InternVL-Performance/resolve/main/internvl3/ocr.png)
111
+
112
+ ### Multi-Image & Real-World Comprehension
113
+
114
+ ![image/png](https://huggingface.co/datasets/Weiyun1025/InternVL-Performance/resolve/main/internvl3/multi-images.png)
115
+
116
+ ### Comprehensive Multimodal & Hallucination Evaluation
117
+
118
+ ![image/png](https://huggingface.co/datasets/Weiyun1025/InternVL-Performance/resolve/main/internvl3/comprehensive.png)
119
+
120
+ ### Visual Grounding
121
+
122
+ ![image/png](https://huggingface.co/datasets/Weiyun1025/InternVL-Performance/resolve/main/internvl3/grounding.png)
123
+
124
+ ### Multimodal Multilingual Understanding
125
+
126
+ ![image/png](https://huggingface.co/datasets/Weiyun1025/InternVL-Performance/resolve/main/internvl3/multilingual.png)
127
+
128
+ ### Video Understanding
129
+
130
+ ![image/png](https://huggingface.co/datasets/Weiyun1025/InternVL-Performance/resolve/main/internvl3/video.png)
131
+
132
+ ### GUI Grounding
133
+
134
+ ![image/png](https://huggingface.co/datasets/Weiyun1025/InternVL-Performance/resolve/main/internvl3/gui.png)
135
+
136
+ ### Spatial Reasoning
137
+
138
+ ![image/png](https://huggingface.co/datasets/Weiyun1025/InternVL-Performance/resolve/main/internvl3/vsi.png)
139
+
140
+ ## Evaluation on Language Capability
141
+
142
+ We compare InternVL3 with Qwen2.5 Chat models, whose corresponding pre-trained base models are employed as the initialization of the langauge component in InternVL3.
143
+ Benefitting from Native Multimodal Pre-Training, the InternVL3 series achieves even better overall text performance than the Qwen2.5 series.
144
+ Please note that the evaluation scores of Qwen2.5 series may differ from those officially reported, as we have adopted the prompt versions provided in the table across all datasets for OpenCompass evaluation.
145
+
146
+ ![image/png](https://huggingface.co/datasets/Weiyun1025/InternVL-Performance/resolve/main/internvl3/text.png)
147
+
148
+ ## Ablation Study
149
+
150
+ ### Native Multimodal Pre-Training
151
+
152
+ We conduct experiments on the InternVL2-8B model while keeping its architecture, initialization parameters, and training data entirely unchanged. Traditionally, InternVL2-8B employs a training pipeline that begins with an MLP warmup phase for feature alignment followed by an Instruction Tuning stage. In our experiments, we substitute the conventional MLP warmup phase with a native multimodal pre-training process. This modification isolates the contribution of native multimodal pre-training to the overall multimodal capability of the model.
153
+
154
+ The evaluation results in the Figure below shows that the model with native multimodal pre-training exhibits performance on most benchmarks that is comparable to the fully multi-stage-trained InternVL2-8B baseline. Furthermore, when followed by instruction tuning on higher-quality data, the model demonstrates further performance gains across evaluated multimodal tasks. These findings underscore the efficiency of native multimodal pre-training in imparting powerful multimodal capabilities to MLLMs.
155
+
156
+ ![image/png](https://huggingface.co/datasets/Weiyun1025/InternVL-Performance/resolve/main/internvl3/ablation-native.png)
157
+
158
+ ### Mixed Preference Optimization
159
+
160
+ As shown in the table below, models fine-tuned with MPO demonstrate superior reasoning performance across seven multimodal reasoning benchmarks compared to their counterparts without MPO. Specifically, InternVL3-78B and InternVL3-38B outperform their counterparts by 4.1 and 4.5 points, respectively. Notably, the training data used for MPO is a subset of that used for SFT, indicating that the performance improvements primarily stem from the training algorithm rather than the training data.
161
+
162
+ ![image/png](https://huggingface.co/datasets/OpenGVLab/MMPR-v1.2-prompts/resolve/main/ablation-mpo.png)
163
+
164
+ ### Variable Visual Position Encoding
165
+
166
+ As reported in the table below, the introduction of V2PE leads to significant performance gains across most evaluation metrics. In addition, our ablation studies—by varying the positional increment \\( \delta \\)—reveal that even for tasks primarily involving conventional contexts, relatively small \\( \delta \\) values can achieve optimal performance. These findings provide important insights for future efforts aimed at refining position encoding strategies for visual tokens in MLLMs.
167
+
168
+ ![image/png](https://huggingface.co/datasets/Weiyun1025/InternVL-Performance/resolve/main/internvl3/ablation-v2pe.png)
169
+
170
+ ## Quick Start
171
+
172
+ We provide an example code to run `InternVL3-9B` using `transformers`.
173
+
174
+ > Please use transformers>=4.37.2 to ensure the model works normally.
175
+
176
+ ### Model Loading
177
+
178
+ #### 16-bit (bf16 / fp16)
179
+
180
+ ```python
181
+ import torch
182
+ from transformers import AutoTokenizer, AutoModel
183
+ path = "OpenGVLab/InternVL3-9B"
184
+ model = AutoModel.from_pretrained(
185
+ path,
186
+ torch_dtype=torch.bfloat16,
187
+ low_cpu_mem_usage=True,
188
+ use_flash_attn=True,
189
+ trust_remote_code=True).eval().cuda()
190
+ ```
191
+
192
+ #### BNB 8-bit Quantization
193
+
194
+ ```python
195
+ import torch
196
+ from transformers import AutoTokenizer, AutoModel
197
+ path = "OpenGVLab/InternVL3-9B"
198
+ model = AutoModel.from_pretrained(
199
+ path,
200
+ torch_dtype=torch.bfloat16,
201
+ load_in_8bit=True,
202
+ low_cpu_mem_usage=True,
203
+ use_flash_attn=True,
204
+ trust_remote_code=True).eval()
205
+ ```
206
+
207
+ #### Multiple GPUs
208
+
209
+ The reason for writing the code this way is to avoid errors that occur during multi-GPU inference due to tensors not being on the same device. By ensuring that the first and last layers of the large language model (LLM) are on the same device, we prevent such errors.
210
+
211
+ ```python
212
+ import math
213
+ import torch
214
+ from transformers import AutoTokenizer, AutoModel
215
+
216
+ def split_model(model_name):
217
+ device_map = {}
218
+ world_size = torch.cuda.device_count()
219
+ config = AutoConfig.from_pretrained(model_path, trust_remote_code=True)
220
+ num_layers = config.llm_config.num_hidden_layers
221
+ # Since the first GPU will be used for ViT, treat it as half a GPU.
222
+ num_layers_per_gpu = math.ceil(num_layers / (world_size - 0.5))
223
+ num_layers_per_gpu = [num_layers_per_gpu] * world_size
224
+ num_layers_per_gpu[0] = math.ceil(num_layers_per_gpu[0] * 0.5)
225
+ layer_cnt = 0
226
+ for i, num_layer in enumerate(num_layers_per_gpu):
227
+ for j in range(num_layer):
228
+ device_map[f'language_model.model.layers.{layer_cnt}'] = i
229
+ layer_cnt += 1
230
+ device_map['vision_model'] = 0
231
+ device_map['mlp1'] = 0
232
+ device_map['language_model.model.tok_embeddings'] = 0
233
+ device_map['language_model.model.embed_tokens'] = 0
234
+ device_map['language_model.output'] = 0
235
+ device_map['language_model.model.norm'] = 0
236
+ device_map['language_model.model.rotary_emb'] = 0
237
+ device_map['language_model.lm_head'] = 0
238
+ device_map[f'language_model.model.layers.{num_layers - 1}'] = 0
239
+
240
+ return device_map
241
+
242
+ path = "OpenGVLab/InternVL3-9B"
243
+ device_map = split_model('InternVL3-9B')
244
+ model = AutoModel.from_pretrained(
245
+ path,
246
+ torch_dtype=torch.bfloat16,
247
+ low_cpu_mem_usage=True,
248
+ use_flash_attn=True,
249
+ trust_remote_code=True,
250
+ device_map=device_map).eval()
251
+ ```
252
+
253
+ ### Inference with Transformers
254
+
255
+ ```python
256
+ import math
257
+ import numpy as np
258
+ import torch
259
+ import torchvision.transforms as T
260
+ from decord import VideoReader, cpu
261
+ from PIL import Image
262
+ from torchvision.transforms.functional import InterpolationMode
263
+ from transformers import AutoModel, AutoTokenizer
264
+
265
+ IMAGENET_MEAN = (0.485, 0.456, 0.406)
266
+ IMAGENET_STD = (0.229, 0.224, 0.225)
267
+
268
+ def build_transform(input_size):
269
+ MEAN, STD = IMAGENET_MEAN, IMAGENET_STD
270
+ transform = T.Compose([
271
+ T.Lambda(lambda img: img.convert('RGB') if img.mode != 'RGB' else img),
272
+ T.Resize((input_size, input_size), interpolation=InterpolationMode.BICUBIC),
273
+ T.ToTensor(),
274
+ T.Normalize(mean=MEAN, std=STD)
275
+ ])
276
+ return transform
277
+
278
+ def find_closest_aspect_ratio(aspect_ratio, target_ratios, width, height, image_size):
279
+ best_ratio_diff = float('inf')
280
+ best_ratio = (1, 1)
281
+ area = width * height
282
+ for ratio in target_ratios:
283
+ target_aspect_ratio = ratio[0] / ratio[1]
284
+ ratio_diff = abs(aspect_ratio - target_aspect_ratio)
285
+ if ratio_diff < best_ratio_diff:
286
+ best_ratio_diff = ratio_diff
287
+ best_ratio = ratio
288
+ elif ratio_diff == best_ratio_diff:
289
+ if area > 0.5 * image_size * image_size * ratio[0] * ratio[1]:
290
+ best_ratio = ratio
291
+ return best_ratio
292
+
293
+ def dynamic_preprocess(image, min_num=1, max_num=12, image_size=448, use_thumbnail=False):
294
+ orig_width, orig_height = image.size
295
+ aspect_ratio = orig_width / orig_height
296
+
297
+ # calculate the existing image aspect ratio
298
+ target_ratios = set(
299
+ (i, j) for n in range(min_num, max_num + 1) for i in range(1, n + 1) for j in range(1, n + 1) if
300
+ i * j <= max_num and i * j >= min_num)
301
+ target_ratios = sorted(target_ratios, key=lambda x: x[0] * x[1])
302
+
303
+ # find the closest aspect ratio to the target
304
+ target_aspect_ratio = find_closest_aspect_ratio(
305
+ aspect_ratio, target_ratios, orig_width, orig_height, image_size)
306
+
307
+ # calculate the target width and height
308
+ target_width = image_size * target_aspect_ratio[0]
309
+ target_height = image_size * target_aspect_ratio[1]
310
+ blocks = target_aspect_ratio[0] * target_aspect_ratio[1]
311
+
312
+ # resize the image
313
+ resized_img = image.resize((target_width, target_height))
314
+ processed_images = []
315
+ for i in range(blocks):
316
+ box = (
317
+ (i % (target_width // image_size)) * image_size,
318
+ (i // (target_width // image_size)) * image_size,
319
+ ((i % (target_width // image_size)) + 1) * image_size,
320
+ ((i // (target_width // image_size)) + 1) * image_size
321
+ )
322
+ # split the image
323
+ split_img = resized_img.crop(box)
324
+ processed_images.append(split_img)
325
+ assert len(processed_images) == blocks
326
+ if use_thumbnail and len(processed_images) != 1:
327
+ thumbnail_img = image.resize((image_size, image_size))
328
+ processed_images.append(thumbnail_img)
329
+ return processed_images
330
+
331
+ def load_image(image_file, input_size=448, max_num=12):
332
+ image = Image.open(image_file).convert('RGB')
333
+ transform = build_transform(input_size=input_size)
334
+ images = dynamic_preprocess(image, image_size=input_size, use_thumbnail=True, max_num=max_num)
335
+ pixel_values = [transform(image) for image in images]
336
+ pixel_values = torch.stack(pixel_values)
337
+ return pixel_values
338
+
339
+ def split_model(model_name):
340
+ device_map = {}
341
+ world_size = torch.cuda.device_count()
342
+ config = AutoConfig.from_pretrained(model_path, trust_remote_code=True)
343
+ num_layers = config.llm_config.num_hidden_layers
344
+ # Since the first GPU will be used for ViT, treat it as half a GPU.
345
+ num_layers_per_gpu = math.ceil(num_layers / (world_size - 0.5))
346
+ num_layers_per_gpu = [num_layers_per_gpu] * world_size
347
+ num_layers_per_gpu[0] = math.ceil(num_layers_per_gpu[0] * 0.5)
348
+ layer_cnt = 0
349
+ for i, num_layer in enumerate(num_layers_per_gpu):
350
+ for j in range(num_layer):
351
+ device_map[f'language_model.model.layers.{layer_cnt}'] = i
352
+ layer_cnt += 1
353
+ device_map['vision_model'] = 0
354
+ device_map['mlp1'] = 0
355
+ device_map['language_model.model.tok_embeddings'] = 0
356
+ device_map['language_model.model.embed_tokens'] = 0
357
+ device_map['language_model.output'] = 0
358
+ device_map['language_model.model.norm'] = 0
359
+ device_map['language_model.model.rotary_emb'] = 0
360
+ device_map['language_model.lm_head'] = 0
361
+ device_map[f'language_model.model.layers.{num_layers - 1}'] = 0
362
+
363
+ return device_map
364
+
365
+ # If you set `load_in_8bit=True`, you will need two 80GB GPUs.
366
+ # If you set `load_in_8bit=False`, you will need at least three 80GB GPUs.
367
+ path = 'OpenGVLab/InternVL3-9B'
368
+ device_map = split_model('InternVL3-9B')
369
+ model = AutoModel.from_pretrained(
370
+ path,
371
+ torch_dtype=torch.bfloat16,
372
+ load_in_8bit=False,
373
+ low_cpu_mem_usage=True,
374
+ use_flash_attn=True,
375
+ trust_remote_code=True,
376
+ device_map=device_map).eval()
377
+ tokenizer = AutoTokenizer.from_pretrained(path, trust_remote_code=True, use_fast=False)
378
+
379
+ # set the max number of tiles in `max_num`
380
+ pixel_values = load_image('./examples/image1.jpg', max_num=12).to(torch.bfloat16).cuda()
381
+ generation_config = dict(max_new_tokens=1024, do_sample=True)
382
+
383
+ # pure-text conversation (纯文本对话)
384
+ question = 'Hello, who are you?'
385
+ response, history = model.chat(tokenizer, None, question, generation_config, history=None, return_history=True)
386
+ print(f'User: {question}\nAssistant: {response}')
387
+
388
+ question = 'Can you tell me a story?'
389
+ response, history = model.chat(tokenizer, None, question, generation_config, history=history, return_history=True)
390
+ print(f'User: {question}\nAssistant: {response}')
391
+
392
+ # single-image single-round conversation (单图单轮对话)
393
+ question = '<image>\nPlease describe the image shortly.'
394
+ response = model.chat(tokenizer, pixel_values, question, generation_config)
395
+ print(f'User: {question}\nAssistant: {response}')
396
+
397
+ # single-image multi-round conversation (单图多轮对话)
398
+ question = '<image>\nPlease describe the image in detail.'
399
+ response, history = model.chat(tokenizer, pixel_values, question, generation_config, history=None, return_history=True)
400
+ print(f'User: {question}\nAssistant: {response}')
401
+
402
+ question = 'Please write a poem according to the image.'
403
+ response, history = model.chat(tokenizer, pixel_values, question, generation_config, history=history, return_history=True)
404
+ print(f'User: {question}\nAssistant: {response}')
405
+
406
+ # multi-image multi-round conversation, combined images (多图多轮对话,拼接图像)
407
+ pixel_values1 = load_image('./examples/image1.jpg', max_num=12).to(torch.bfloat16).cuda()
408
+ pixel_values2 = load_image('./examples/image2.jpg', max_num=12).to(torch.bfloat16).cuda()
409
+ pixel_values = torch.cat((pixel_values1, pixel_values2), dim=0)
410
+
411
+ question = '<image>\nDescribe the two images in detail.'
412
+ response, history = model.chat(tokenizer, pixel_values, question, generation_config,
413
+ history=None, return_history=True)
414
+ print(f'User: {question}\nAssistant: {response}')
415
+
416
+ question = 'What are the similarities and differences between these two images.'
417
+ response, history = model.chat(tokenizer, pixel_values, question, generation_config,
418
+ history=history, return_history=True)
419
+ print(f'User: {question}\nAssistant: {response}')
420
+
421
+ # multi-image multi-round conversation, separate images (多图多轮对话,独立图像)
422
+ pixel_values1 = load_image('./examples/image1.jpg', max_num=12).to(torch.bfloat16).cuda()
423
+ pixel_values2 = load_image('./examples/image2.jpg', max_num=12).to(torch.bfloat16).cuda()
424
+ pixel_values = torch.cat((pixel_values1, pixel_values2), dim=0)
425
+ num_patches_list = [pixel_values1.size(0), pixel_values2.size(0)]
426
+
427
+ question = 'Image-1: <image>\nImage-2: <image>\nDescribe the two images in detail.'
428
+ response, history = model.chat(tokenizer, pixel_values, question, generation_config,
429
+ num_patches_list=num_patches_list,
430
+ history=None, return_history=True)
431
+ print(f'User: {question}\nAssistant: {response}')
432
+
433
+ question = 'What are the similarities and differences between these two images.'
434
+ response, history = model.chat(tokenizer, pixel_values, question, generation_config,
435
+ num_patches_list=num_patches_list,
436
+ history=history, return_history=True)
437
+ print(f'User: {question}\nAssistant: {response}')
438
+
439
+ # batch inference, single image per sample (单图批处理)
440
+ pixel_values1 = load_image('./examples/image1.jpg', max_num=12).to(torch.bfloat16).cuda()
441
+ pixel_values2 = load_image('./examples/image2.jpg', max_num=12).to(torch.bfloat16).cuda()
442
+ num_patches_list = [pixel_values1.size(0), pixel_values2.size(0)]
443
+ pixel_values = torch.cat((pixel_values1, pixel_values2), dim=0)
444
+
445
+ questions = ['<image>\nDescribe the image in detail.'] * len(num_patches_list)
446
+ responses = model.batch_chat(tokenizer, pixel_values,
447
+ num_patches_list=num_patches_list,
448
+ questions=questions,
449
+ generation_config=generation_config)
450
+ for question, response in zip(questions, responses):
451
+ print(f'User: {question}\nAssistant: {response}')
452
+
453
+ # video multi-round conversation (视频多轮对话)
454
+ def get_index(bound, fps, max_frame, first_idx=0, num_segments=32):
455
+ if bound:
456
+ start, end = bound[0], bound[1]
457
+ else:
458
+ start, end = -100000, 100000
459
+ start_idx = max(first_idx, round(start * fps))
460
+ end_idx = min(round(end * fps), max_frame)
461
+ seg_size = float(end_idx - start_idx) / num_segments
462
+ frame_indices = np.array([
463
+ int(start_idx + (seg_size / 2) + np.round(seg_size * idx))
464
+ for idx in range(num_segments)
465
+ ])
466
+ return frame_indices
467
+
468
+ def load_video(video_path, bound=None, input_size=448, max_num=1, num_segments=32):
469
+ vr = VideoReader(video_path, ctx=cpu(0), num_threads=1)
470
+ max_frame = len(vr) - 1
471
+ fps = float(vr.get_avg_fps())
472
+
473
+ pixel_values_list, num_patches_list = [], []
474
+ transform = build_transform(input_size=input_size)
475
+ frame_indices = get_index(bound, fps, max_frame, first_idx=0, num_segments=num_segments)
476
+ for frame_index in frame_indices:
477
+ img = Image.fromarray(vr[frame_index].asnumpy()).convert('RGB')
478
+ img = dynamic_preprocess(img, image_size=input_size, use_thumbnail=True, max_num=max_num)
479
+ pixel_values = [transform(tile) for tile in img]
480
+ pixel_values = torch.stack(pixel_values)
481
+ num_patches_list.append(pixel_values.shape[0])
482
+ pixel_values_list.append(pixel_values)
483
+ pixel_values = torch.cat(pixel_values_list)
484
+ return pixel_values, num_patches_list
485
+
486
+ video_path = './examples/red-panda.mp4'
487
+ pixel_values, num_patches_list = load_video(video_path, num_segments=8, max_num=1)
488
+ pixel_values = pixel_values.to(torch.bfloat16).cuda()
489
+ video_prefix = ''.join([f'Frame{i+1}: <image>\n' for i in range(len(num_patches_list))])
490
+ question = video_prefix + 'What is the red panda doing?'
491
+ # Frame1: <image>\nFrame2: <image>\n...\nFrame8: <image>\n{question}
492
+ response, history = model.chat(tokenizer, pixel_values, question, generation_config,
493
+ num_patches_list=num_patches_list, history=None, return_history=True)
494
+ print(f'User: {question}\nAssistant: {response}')
495
+
496
+ question = 'Describe this video in detail.'
497
+ response, history = model.chat(tokenizer, pixel_values, question, generation_config,
498
+ num_patches_list=num_patches_list, history=history, return_history=True)
499
+ print(f'User: {question}\nAssistant: {response}')
500
+ ```
501
+
502
+ #### Streaming Output
503
+
504
+ Besides this method, you can also use the following code to get streamed output.
505
+
506
+ ```python
507
+ from transformers import TextIteratorStreamer
508
+ from threading import Thread
509
+
510
+ # Initialize the streamer
511
+ streamer = TextIteratorStreamer(tokenizer, skip_prompt=True, skip_special_tokens=True, timeout=10)
512
+ # Define the generation configuration
513
+ generation_config = dict(max_new_tokens=1024, do_sample=False, streamer=streamer)
514
+ # Start the model chat in a separate thread
515
+ thread = Thread(target=model.chat, kwargs=dict(
516
+ tokenizer=tokenizer, pixel_values=pixel_values, question=question,
517
+ history=None, return_history=False, generation_config=generation_config,
518
+ ))
519
+ thread.start()
520
+
521
+ # Initialize an empty string to store the generated text
522
+ generated_text = ''
523
+ # Loop through the streamer to get the new text as it is generated
524
+ for new_text in streamer:
525
+ if new_text == model.conv_template.sep:
526
+ break
527
+ generated_text += new_text
528
+ print(new_text, end='', flush=True) # Print each new chunk of generated text on the same line
529
+ ```
530
+
531
+ ## Finetune
532
+
533
+ Many repositories now support fine-tuning of the InternVL series models, including [InternVL](https://github.com/OpenGVLab/InternVL), [SWIFT](https://github.com/modelscope/ms-swift), [XTurner](https://github.com/InternLM/xtuner), and others. Please refer to their documentation for more details on fine-tuning.
534
+
535
+ ## Deployment
536
+
537
+ ### LMDeploy
538
+
539
+ LMDeploy is a toolkit for compressing, deploying, and serving LLMs & VLMs.
540
+
541
+ ```sh
542
+ # if lmdeploy<0.7.3, you need to explicitly set chat_template_config=ChatTemplateConfig(model_name='internvl2_5')
543
+ pip install lmdeploy>=0.7.3
544
+ ```
545
+
546
+ LMDeploy abstracts the complex inference process of multi-modal Vision-Language Models (VLM) into an easy-to-use pipeline, similar to the Large Language Model (LLM) inference pipeline.
547
+
548
+ #### A 'Hello, world' Example
549
+
550
+ ```python
551
+ from lmdeploy import pipeline, TurbomindEngineConfig, ChatTemplateConfig
552
+ from lmdeploy.vl import load_image
553
+
554
+ model = 'OpenGVLab/InternVL3-9B'
555
+ image = load_image('https://raw.githubusercontent.com/open-mmlab/mmdeploy/main/tests/data/tiger.jpeg')
556
+ pipe = pipeline(model, backend_config=TurbomindEngineConfig(session_len=16384, tp=1), chat_template_config=ChatTemplateConfig(model_name='internvl2_5'))
557
+ response = pipe(('describe this image', image))
558
+ print(response.text)
559
+ ```
560
+
561
+ If `ImportError` occurs while executing this case, please install the required dependency packages as prompted.
562
+
563
+ #### Multi-images Inference
564
+
565
+ When dealing with multiple images, you can put them all in one list. Keep in mind that multiple images will lead to a higher number of input tokens, and as a result, the size of the context window typically needs to be increased.
566
+
567
+ ```python
568
+ from lmdeploy import pipeline, TurbomindEngineConfig, ChatTemplateConfig
569
+ from lmdeploy.vl import load_image
570
+ from lmdeploy.vl.constants import IMAGE_TOKEN
571
+
572
+ model = 'OpenGVLab/InternVL3-9B'
573
+ pipe = pipeline(model, backend_config=TurbomindEngineConfig(session_len=16384, tp=1), chat_template_config=ChatTemplateConfig(model_name='internvl2_5'))
574
+
575
+ image_urls=[
576
+ 'https://raw.githubusercontent.com/open-mmlab/mmdeploy/main/demo/resources/human-pose.jpg',
577
+ 'https://raw.githubusercontent.com/open-mmlab/mmdeploy/main/demo/resources/det.jpg'
578
+ ]
579
+
580
+ images = [load_image(img_url) for img_url in image_urls]
581
+ # Numbering images improves multi-image conversations
582
+ response = pipe((f'Image-1: {IMAGE_TOKEN}\nImage-2: {IMAGE_TOKEN}\ndescribe these two images', images))
583
+ print(response.text)
584
+ ```
585
+
586
+ #### Batch Prompts Inference
587
+
588
+ Conducting inference with batch prompts is quite straightforward; just place them within a list structure:
589
+
590
+ ```python
591
+ from lmdeploy import pipeline, TurbomindEngineConfig, ChatTemplateConfig
592
+ from lmdeploy.vl import load_image
593
+
594
+ model = 'OpenGVLab/InternVL3-9B'
595
+ pipe = pipeline(model, backend_config=TurbomindEngineConfig(session_len=16384, tp=1), chat_template_config=ChatTemplateConfig(model_name='internvl2_5'))
596
+
597
+ image_urls=[
598
+ "https://raw.githubusercontent.com/open-mmlab/mmdeploy/main/demo/resources/human-pose.jpg",
599
+ "https://raw.githubusercontent.com/open-mmlab/mmdeploy/main/demo/resources/det.jpg"
600
+ ]
601
+ prompts = [('describe this image', load_image(img_url)) for img_url in image_urls]
602
+ response = pipe(prompts)
603
+ print(response)
604
+ ```
605
+
606
+ #### Multi-turn Conversation
607
+
608
+ There are two ways to do the multi-turn conversations with the pipeline. One is to construct messages according to the format of OpenAI and use above introduced method, the other is to use the `pipeline.chat` interface.
609
+
610
+ ```python
611
+ from lmdeploy import pipeline, TurbomindEngineConfig, GenerationConfig, ChatTemplateConfig
612
+ from lmdeploy.vl import load_image
613
+
614
+ model = 'OpenGVLab/InternVL3-9B'
615
+ pipe = pipeline(model, backend_config=TurbomindEngineConfig(session_len=16384, tp=1), chat_template_config=ChatTemplateConfig(model_name='internvl2_5'))
616
+
617
+ image = load_image('https://raw.githubusercontent.com/open-mmlab/mmdeploy/main/demo/resources/human-pose.jpg')
618
+ gen_config = GenerationConfig(top_k=40, top_p=0.8, temperature=0.8)
619
+ sess = pipe.chat(('describe this image', image), gen_config=gen_config)
620
+ print(sess.response.text)
621
+ sess = pipe.chat('What is the woman doing?', session=sess, gen_config=gen_config)
622
+ print(sess.response.text)
623
+ ```
624
+
625
+ #### Service
626
+
627
+ LMDeploy's `api_server` enables models to be easily packed into services with a single command. The provided RESTful APIs are compatible with OpenAI's interfaces. Below are an example of service startup:
628
+
629
+ ```shell
630
+ lmdeploy serve api_server OpenGVLab/InternVL3-9B --chat-template internvl2_5 --server-port 23333 --tp 1
631
+ ```
632
+
633
+ To use the OpenAI-style interface, you need to install OpenAI:
634
+
635
+ ```shell
636
+ pip install openai
637
+ ```
638
+
639
+ Then, use the code below to make the API call:
640
+
641
+ ```python
642
+ from openai import OpenAI
643
+
644
+ client = OpenAI(api_key='YOUR_API_KEY', base_url='http://0.0.0.0:23333/v1')
645
+ model_name = client.models.list().data[0].id
646
+ response = client.chat.completions.create(
647
+ model=model_name,
648
+ messages=[{
649
+ 'role':
650
+ 'user',
651
+ 'content': [{
652
+ 'type': 'text',
653
+ 'text': 'describe this image',
654
+ }, {
655
+ 'type': 'image_url',
656
+ 'image_url': {
657
+ 'url':
658
+ 'https://modelscope.oss-cn-beijing.aliyuncs.com/resource/tiger.jpeg',
659
+ },
660
+ }],
661
+ }],
662
+ temperature=0.8,
663
+ top_p=0.8)
664
+ print(response)
665
+ ```
666
+
667
+ ## License
668
+
669
+ This project is released under the MIT License.
670
+
671
+ ## Citation
672
+
673
+ If you find this project useful in your research, please consider citing:
674
+
675
+ ```BibTeX
676
+ @article{chen2024expanding,
677
+ title={Expanding Performance Boundaries of Open-Source Multimodal Models with Model, Data, and Test-Time Scaling},
678
+ author={Chen, Zhe and Wang, Weiyun and Cao, Yue and Liu, Yangzhou and Gao, Zhangwei and Cui, Erfei and Zhu, Jinguo and Ye, Shenglong and Tian, Hao and Liu, Zhaoyang and others},
679
+ journal={arXiv preprint arXiv:2412.05271},
680
+ year={2024}
681
+ }
682
+ @article{wang2024mpo,
683
+ title={Enhancing the Reasoning Ability of Multimodal Large Language Models via Mixed Preference Optimization},
684
+ author={Wang, Weiyun and Chen, Zhe and Wang, Wenhai and Cao, Yue and Liu, Yangzhou and Gao, Zhangwei and Zhu, Jinguo and Zhu, Xizhou and Lu, Lewei and Qiao, Yu and Dai, Jifeng},
685
+ journal={arXiv preprint arXiv:2411.10442},
686
+ year={2024}
687
+ }
688
+ @article{chen2024far,
689
+ title={How Far Are We to GPT-4V? Closing the Gap to Commercial Multimodal Models with Open-Source Suites},
690
+ author={Chen, Zhe and Wang, Weiyun and Tian, Hao and Ye, Shenglong and Gao, Zhangwei and Cui, Erfei and Tong, Wenwen and Hu, Kongzhi and Luo, Jiapeng and Ma, Zheng and others},
691
+ journal={arXiv preprint arXiv:2404.16821},
692
+ year={2024}
693
+ }
694
+ @inproceedings{chen2024internvl,
695
+ title={Internvl: Scaling up vision foundation models and aligning for generic visual-linguistic tasks},
696
+ author={Chen, Zhe and Wu, Jiannan and Wang, Wenhai and Su, Weijie and Chen, Guo and Xing, Sen and Zhong, Muyan and Zhang, Qinglong and Zhu, Xizhou and Lu, Lewei and others},
697
+ booktitle={Proceedings of the IEEE/CVF Conference on Computer Vision and Pattern Recognition},
698
+ pages={24185--24198},
699
+ year={2024}
700
+ }
701
+ ```
merged_transition.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:4189bcb7300b89ff02c1631ec8644b6fce8d119e1dee409336b54404cb203b27
3
+ size 645923336
modeling_cvrr_merged.py ADDED
@@ -0,0 +1,279 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Inference-only merged CVRR, preserving the reference strict decoding path.
2
+
3
+ Native projections and dense learned projections are separate. No A/B adapter
4
+ tensors or external backbone download is required after loading the package.
5
+ """
6
+ from __future__ import annotations
7
+
8
+ import contextlib
9
+ import json
10
+ import types
11
+ from pathlib import Path
12
+ from typing import TYPE_CHECKING
13
+
14
+ import torch
15
+ from torch import nn
16
+ from safetensors.torch import load_file
17
+ from transformers import PreTrainedModel
18
+
19
+ from .configuration_cvrr_merged import CVRRMergedConfig
20
+
21
+ # Transformers 4.57's local-directory loader copies direct relative imports;
22
+ # enumerate transitive helpers without eagerly importing every backbone.
23
+ if TYPE_CHECKING:
24
+ from .source_splitting import SplitContext
25
+ from .source_perceive import PerceiveDeliberateLayout
26
+ from .source_spatial import SpatialVisualRecurrentCell
27
+ from .source_helpers import _layer_hidden
28
+ from .source_gemma import GemmaCVRR
29
+ from .source_internvl import InternVLCVRR
30
+ from .configuration_source_qwen25 import CloseQwen2_5_VLConfig
31
+ from .configuration_source_qwen3 import CloseQwen3VLConfig
32
+ from .modeling_source_qwen25 import CloseQwen2_5_VLForConditionalGeneration
33
+ from .modeling_source_qwen3 import CloseQwen3VLForConditionalGeneration
34
+
35
+
36
+ class NativeMergedLinear(nn.Module):
37
+ def __init__(self, native, merged):
38
+ super().__init__()
39
+ if tuple(native.weight.shape) != tuple(merged.shape):
40
+ raise ValueError('Native/merged projection shape mismatch')
41
+ self.base = native
42
+ self.register_buffer('merged_weight', merged.to(device=native.weight.device, dtype=torch.float32))
43
+ self.enabled = False
44
+ self.dropout = nn.Identity()
45
+ self.requires_grad_(False)
46
+
47
+ def _apply(self, fn, recurse=True):
48
+ # .to(device) is allowed; a caller's dtype cast must not truncate the
49
+ # FP32 merged weights and then merely upcast already-lost precision.
50
+ original = self.merged_weight
51
+ super()._apply(fn, recurse=recurse)
52
+ self.merged_weight = original.to(device=self.merged_weight.device, dtype=torch.float32)
53
+ return self
54
+
55
+ def forward(self, x):
56
+ if not self.enabled:
57
+ return self.base(x)
58
+ bias = None if self.base.bias is None else self.base.bias.float()
59
+ return torch.nn.functional.linear(x.float(), self.merged_weight, bias).to(x.dtype)
60
+
61
+
62
+ @contextlib.contextmanager
63
+ def merged_execution(model, enabled):
64
+ modules = model._cvrr_merged_projections
65
+ previous = [m.enabled for m in modules]
66
+ for m in modules:
67
+ m.enabled = bool(enabled)
68
+ try:
69
+ yield
70
+ finally:
71
+ for m, old in zip(modules, previous):
72
+ m.enabled = old
73
+
74
+
75
+ def set_merged_execution(model, enabled):
76
+ for m in model._cvrr_merged_projections:
77
+ m.enabled = bool(enabled)
78
+
79
+
80
+ def install_merged(runtime, weights, *, generic):
81
+ cell = runtime.layers[runtime.cell_index] if generic else runtime.text_model.layers[int(runtime.config.ell_star) + 1]
82
+ modules = []
83
+ for key, weight in sorted(weights.items()):
84
+ if not key.endswith('.weight'):
85
+ raise ValueError(f'Unexpected merged state key: {key}')
86
+ parts = key.removesuffix('.weight').split('.')
87
+ parent = cell
88
+ for part in parts[:-1]:
89
+ parent = getattr(parent, part)
90
+ old = getattr(parent, parts[-1])
91
+ native = old.base if generic else old.base_layer
92
+ wrapped = NativeMergedLinear(native, weight)
93
+ setattr(parent, parts[-1], wrapped)
94
+ modules.append(wrapped)
95
+ # Deliberately a tuple, not another ModuleList alias in state_dict.
96
+ runtime._cvrr_merged_projections = tuple(modules)
97
+ if generic:
98
+ runtime.lora = {str(i): m for i, m in enumerate(modules)}
99
+ runtime.adapters = types.MethodType(merged_execution, runtime)
100
+ else:
101
+ runtime._recurrence_adapter_execution = types.MethodType(merged_execution, runtime)
102
+ runtime._set_adapters_enabled = types.MethodType(set_merged_execution, runtime)
103
+ if any('lora_' in n for n, _ in runtime.named_parameters()):
104
+ raise RuntimeError('Unmerged LoRA parameters remain')
105
+ runtime.requires_grad_(False)
106
+ return runtime.eval()
107
+
108
+
109
+ class CVRRMergedModel(PreTrainedModel):
110
+ config_class = CVRRMergedConfig
111
+ main_input_name = 'input_ids'
112
+
113
+ def __init__(self, config, runtime=None):
114
+ super().__init__(config)
115
+ if runtime is None:
116
+ raise ValueError('Load a packaged release using from_pretrained')
117
+ self.runtime = runtime
118
+ self.is_generic = 'Qwen' not in config.release['name']
119
+
120
+ @classmethod
121
+ def from_pretrained(cls, pretrained_model_name_or_path, *args, **kwargs):
122
+ if args:
123
+ raise TypeError('Positional loader overrides are not supported')
124
+ config = kwargs.pop('config', None)
125
+ device = kwargs.pop('device_map', kwargs.pop('device', 'cpu'))
126
+ if device is None:
127
+ device = 'cpu'
128
+ if isinstance(device, dict) or str(device) == 'auto':
129
+ raise ValueError('Use one complete replica per device; pass device_map="cuda:0" or "cpu"')
130
+ dtype = kwargs.pop('dtype', kwargs.pop('torch_dtype', torch.bfloat16))
131
+ if dtype not in (torch.bfloat16, 'bfloat16', 'auto', None):
132
+ raise ValueError('This release preserves BF16 native weights and FP32 merged projections')
133
+ offline = kwargs.pop('local_files_only', False)
134
+ revision = kwargs.pop('revision', None)
135
+ token = kwargs.pop('token', None)
136
+ kwargs.pop('trust_remote_code', None)
137
+ kwargs.pop('_from_auto', None)
138
+ kwargs.pop('adapter_kwargs', None)
139
+ kwargs.pop('name_or_path', None)
140
+ kwargs.pop('_commit_hash', None)
141
+ if kwargs:
142
+ raise TypeError(f'Unsupported loader options: {sorted(kwargs)}')
143
+ path = Path(pretrained_model_name_or_path)
144
+ if not path.is_dir():
145
+ from huggingface_hub import snapshot_download
146
+ path = Path(snapshot_download(str(pretrained_model_name_or_path), revision=revision,
147
+ token=token, local_files_only=offline))
148
+ if config is None:
149
+ config = CVRRMergedConfig.from_pretrained(path, local_files_only=True)
150
+ settings = json.loads((path / 'cvrr_release_config.json').read_text())
151
+ if config.release != settings:
152
+ raise ValueError('Root configuration and release metadata disagree')
153
+ if settings['format'] != 'cvrr_native_plus_merged_transition_v1':
154
+ raise ValueError('Unsupported release layout')
155
+ native_path = str(path / 'native_backbone')
156
+ generic = 'Qwen' not in settings['name']
157
+ common = dict(ell_star=settings['ell_star'], steps=settings['inference_T'],
158
+ beta=settings['inference_beta'], rank=settings['lora_rank'],
159
+ alpha=settings['lora_alpha'], dropout=settings['lora_dropout'],
160
+ device=device, offline=True)
161
+ if generic:
162
+ if 'InternVL' in settings['name']:
163
+ from .source_internvl import InternVLCVRR
164
+ runtime = InternVLCVRR(native_path, **common)
165
+ else:
166
+ from .source_gemma import GemmaCVRR
167
+ runtime = GemmaCVRR(native_path, **common)
168
+ else:
169
+ if 'Qwen2.5' in settings['name']:
170
+ from transformers import Qwen2_5_VLForConditionalGeneration as Native
171
+ from .configuration_source_qwen25 import CloseQwen2_5_VLConfig as SourceConfig
172
+ from .modeling_source_qwen25 import CloseQwen2_5_VLForConditionalGeneration as Source
173
+ else:
174
+ from transformers import Qwen3VLForConditionalGeneration as Native
175
+ from .configuration_source_qwen3 import CloseQwen3VLConfig as SourceConfig
176
+ from .modeling_source_qwen3 import CloseQwen3VLForConditionalGeneration as Source
177
+ source_config = SourceConfig.from_pretrained(path / 'source_config', local_files_only=True)
178
+ source_config.num_workspace_steps = settings['inference_T']
179
+ source_config.counterfactual_beta = settings['inference_beta']
180
+ for flag in ('counterfactual_reliance_loss', 'counterfactual_step_retention_loss',
181
+ 'functional_state_loss', 'interface_loss'):
182
+ setattr(source_config, flag, False)
183
+ native = Native.from_pretrained(native_path, dtype=torch.bfloat16,
184
+ device_map=str(device), local_files_only=True,
185
+ attn_implementation='sdpa')
186
+ runtime = Source(source_config, backbone=native).to(device)
187
+ weights = load_file(path / 'merged_transition.safetensors', device='cpu')
188
+ runtime = install_merged(runtime, weights, generic=generic)
189
+ model = cls(config, runtime).eval()
190
+ model.release_path = path
191
+ return model
192
+
193
+ def prepare_inputs(self, image, question, *, max_visual_tokens=8192, max_tiles=12):
194
+ """Single-image, batch-one preparation using the reference model prompts."""
195
+ if max_visual_tokens < 1:
196
+ raise ValueError('max_visual_tokens must be positive')
197
+ device = next(self.runtime.parameters()).device
198
+ if 'InternVL' in self.config.release['name']:
199
+ from .source_internvl import InternVLArrowCollator
200
+ from .source_helpers import dynamic_tiles, _normalize_tiles
201
+ helper = InternVLArrowCollator(self.runtime, max_tiles=max_tiles)
202
+ tiles = dynamic_tiles(image, image_size=helper.image_size, max_tiles=max_tiles,
203
+ thumbnail=helper.use_thumbnail)
204
+ prompt = helper._query(question, '', len(tiles))
205
+ self.tokenizer = self.runtime.tokenizer
206
+ encoded = dict(self.tokenizer(prompt, return_tensors='pt'))
207
+ encoded.update(pixel_values=_normalize_tiles(tiles).to(torch.bfloat16),
208
+ image_flags=torch.ones(len(tiles), 1, dtype=torch.long))
209
+ else:
210
+ from transformers import AutoProcessor
211
+ processor = AutoProcessor.from_pretrained(self.release_path / 'native_backbone',
212
+ local_files_only=True, use_fast=('Qwen3' in self.config.release['name']))
213
+ self.tokenizer = processor.tokenizer
214
+ if not self.is_generic:
215
+ cfg = self.runtime.config
216
+ patch = int(cfg.vision_config.patch_size) * int(cfg.vision_config.spatial_merge_size)
217
+ pixels = max_visual_tokens * patch * patch
218
+ ip = processor.image_processor
219
+ if isinstance(getattr(ip, 'size', None), dict):
220
+ ip.size = dict(ip.size, longest_edge=pixels)
221
+ ip.max_pixels = pixels
222
+ prompt = processor.apply_chat_template([{'role':'user', 'content':[
223
+ {'type':'image'}, {'type':'text','text':question}]}],
224
+ tokenize=False, add_generation_prompt=True)
225
+ encoded = dict(processor(text=[prompt], images=[image], return_tensors='pt'))
226
+ if not self.is_generic:
227
+ txt = processor.apply_chat_template([{'role':'user','content':[
228
+ {'type':'text','text':question}]}], tokenize=False, add_generation_prompt=True)
229
+ q = processor.tokenizer(txt, return_tensors='pt', add_special_tokens=False)
230
+ encoded['question_ids'] = q['input_ids']
231
+ encoded['question_attention_mask'] = q['attention_mask']
232
+ result = {k: v.to(device) for k,v in encoded.items() if isinstance(v, torch.Tensor)}
233
+ if 'pixel_values' in result:
234
+ result['pixel_values'] = result['pixel_values'].to(torch.bfloat16)
235
+ if self.is_generic:
236
+ count = int((self.runtime._modality(result).eq(1) & result['attention_mask'].bool()).sum())
237
+ else:
238
+ count = int(result['input_ids'].eq(self.runtime.config.image_token_id).sum())
239
+ allowed = {'input_ids','attention_mask','pixel_values','image_grid_thw',
240
+ 'question_ids','question_attention_mask'}
241
+ result = {k:v for k,v in result.items() if k in allowed}
242
+ if not 0 < count <= max_visual_tokens:
243
+ raise ValueError(f'Visual token count {count} exceeds cap {max_visual_tokens} or is empty')
244
+ return result
245
+
246
+ def get_input_embeddings(self):
247
+ if self.is_generic:
248
+ return self.runtime.base_model.get_input_embeddings()
249
+ return self.runtime.backbone.get_input_embeddings()
250
+
251
+ def save_pretrained(self, *args, **kwargs):
252
+ raise NotImplementedError('Keep the exported directory intact; generic HF reserialization would change the release layout')
253
+
254
+ def forward(self, *args, **kwargs):
255
+ raise NotImplementedError('Inference-only release: use next_token_logits or generate')
256
+
257
+ @torch.inference_mode()
258
+ def next_token_logits(self, **inputs):
259
+ if self.is_generic:
260
+ trace = self.runtime.extract(inputs)
261
+ state = self.runtime.rollout(trace)[-1]
262
+ return self.runtime.next_token_logits(trace, state)
263
+ upper, _cache, _states, _evidence = self.runtime.prefill(**inputs)
264
+ mask = inputs.get('question_attention_mask')
265
+ last = (mask.long().sum(-1) - 1 if mask is not None else
266
+ torch.full((upper.shape[0],), upper.shape[1]-1, device=upper.device))
267
+ hidden = upper[torch.arange(upper.shape[0], device=upper.device), last.long()]
268
+ return self.runtime.backbone.lm_head(self.runtime.text_model.norm(hidden)).float()
269
+
270
+ @torch.inference_mode()
271
+ def generate(self, *, do_sample=False, max_new_tokens=32, **inputs):
272
+ if do_sample:
273
+ raise ValueError('Only deterministic greedy decoding is supported by this release')
274
+ if self.is_generic:
275
+ raise NotImplementedError('Gemma/InternVL release validation currently covers strict next-token readout only')
276
+ return self.runtime.generate_answer(**inputs, max_new_tokens=max_new_tokens)
277
+
278
+
279
+ CVRRMergedModel.register_for_auto_class('AutoModelForImageTextToText')
modeling_source_qwen25.py ADDED
The diff for this file is too large to render. See raw diff
 
modeling_source_qwen3.py ADDED
The diff for this file is too large to render. See raw diff
 
native_backbone/added_tokens.json ADDED
@@ -0,0 +1,11 @@
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "</box>": 128141,
3
+ "</img>": 128134,
4
+ "</quad>": 128137,
5
+ "</ref>": 128139,
6
+ "<IMG_CONTEXT>": 128135,
7
+ "<box>": 128140,
8
+ "<img>": 128133,
9
+ "<quad>": 128136,
10
+ "<ref>": 128138
11
+ }
native_backbone/config.json ADDED
@@ -0,0 +1,225 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "_commit_hash": null,
3
+ "architectures": [
4
+ "InternVLChatModel"
5
+ ],
6
+ "auto_map": {
7
+ "AutoConfig": "configuration_internvl_chat.InternVLChatConfig",
8
+ "AutoModel": "modeling_internvl_chat.InternVLChatModel",
9
+ "AutoModelForCausalLM": "modeling_internvl_chat.InternVLChatModel"
10
+ },
11
+ "downsample_ratio": 0.5,
12
+ "dynamic_image_size": true,
13
+ "force_image_size": 448,
14
+ "hidden_size": 4096,
15
+ "image_fold": null,
16
+ "llm_config": {
17
+ "_attn_implementation_autoset": true,
18
+ "add_cross_attention": false,
19
+ "architectures": [
20
+ "InternLM2ForCausalLM"
21
+ ],
22
+ "attn_implementation": "flash_attention_2",
23
+ "auto_map": {
24
+ "AutoConfig": "configuration_internlm2.InternLM2Config",
25
+ "AutoModel": "modeling_internlm2.InternLM2ForCausalLM",
26
+ "AutoModelForCausalLM": "modeling_internlm2.InternLM2ForCausalLM"
27
+ },
28
+ "bad_words_ids": null,
29
+ "begin_suppress_tokens": null,
30
+ "bias": false,
31
+ "bos_token_id": 1,
32
+ "chunk_size_feed_forward": 0,
33
+ "cross_attention_hidden_size": null,
34
+ "decoder_start_token_id": null,
35
+ "diversity_penalty": 0.0,
36
+ "do_sample": false,
37
+ "early_stopping": false,
38
+ "encoder_no_repeat_ngram_size": 0,
39
+ "eos_token_id": 2,
40
+ "exponential_decay_length_penalty": null,
41
+ "finetuning_task": null,
42
+ "forced_bos_token_id": null,
43
+ "forced_eos_token_id": null,
44
+ "hidden_act": "silu",
45
+ "hidden_size": 4096,
46
+ "id2label": {
47
+ "0": "LABEL_0",
48
+ "1": "LABEL_1"
49
+ },
50
+ "initializer_range": 0.02,
51
+ "intermediate_size": 10240,
52
+ "is_decoder": false,
53
+ "is_encoder_decoder": false,
54
+ "label2id": {
55
+ "LABEL_0": 0,
56
+ "LABEL_1": 1
57
+ },
58
+ "length_penalty": 1.0,
59
+ "max_length": 20,
60
+ "max_position_embeddings": 32768,
61
+ "min_length": 0,
62
+ "model_type": "internlm2",
63
+ "moe_config": null,
64
+ "no_repeat_ngram_size": 0,
65
+ "num_attention_heads": 32,
66
+ "num_beam_groups": 1,
67
+ "num_beams": 1,
68
+ "num_hidden_layers": 48,
69
+ "num_key_value_heads": 2,
70
+ "num_return_sequences": 1,
71
+ "output_attentions": false,
72
+ "output_hidden_states": false,
73
+ "output_scores": false,
74
+ "pad_token_id": 2,
75
+ "prefix": null,
76
+ "pretraining_tp": 1,
77
+ "problem_type": null,
78
+ "pruned_heads": {},
79
+ "remove_invalid_values": false,
80
+ "repetition_penalty": 1.0,
81
+ "return_dict": true,
82
+ "return_dict_in_generate": false,
83
+ "rms_norm_eps": 1e-05,
84
+ "rope_scaling": {
85
+ "factor": 2.0,
86
+ "type": "dynamic"
87
+ },
88
+ "rope_theta": 50000000,
89
+ "sep_token_id": null,
90
+ "suppress_tokens": null,
91
+ "task_specific_params": null,
92
+ "temperature": 1.0,
93
+ "tf_legacy_loss": false,
94
+ "tie_encoder_decoder": false,
95
+ "tie_word_embeddings": false,
96
+ "tokenizer_class": null,
97
+ "top_k": 50,
98
+ "top_p": 1.0,
99
+ "torch_dtype": "bfloat16",
100
+ "torchscript": false,
101
+ "transformers_version": "4.48.3",
102
+ "typical_p": 1.0,
103
+ "use_bfloat16": false,
104
+ "use_cache": false,
105
+ "vocab_size": 128142
106
+ },
107
+ "max_dynamic_patch": 12,
108
+ "min_dynamic_patch": 1,
109
+ "model_type": "internvl_chat",
110
+ "pad2square": false,
111
+ "ps_version": "v2",
112
+ "select_layer": -1,
113
+ "system_message": null,
114
+ "template": "internvl2_5",
115
+ "tie_word_embeddings": false,
116
+ "torch_dtype": "bfloat16",
117
+ "transformers_version": null,
118
+ "use_backbone_lora": 0,
119
+ "use_img_start_end_token": true,
120
+ "use_llm_lora": 0,
121
+ "use_thumbnail": true,
122
+ "vision_config": {
123
+ "_attn_implementation_autoset": true,
124
+ "add_cross_attention": false,
125
+ "architectures": [
126
+ "InternVisionModel"
127
+ ],
128
+ "attention_dropout": 0.0,
129
+ "auto_map": {
130
+ "AutoConfig": "configuration_intern_vit.InternVisionConfig",
131
+ "AutoModel": "modeling_intern_vit.InternVisionModel"
132
+ },
133
+ "bad_words_ids": null,
134
+ "begin_suppress_tokens": null,
135
+ "bos_token_id": null,
136
+ "capacity_factor": 1.2,
137
+ "chunk_size_feed_forward": 0,
138
+ "cross_attention_hidden_size": null,
139
+ "decoder_start_token_id": null,
140
+ "diversity_penalty": 0.0,
141
+ "do_sample": false,
142
+ "drop_path_rate": 0.1,
143
+ "dropout": 0.0,
144
+ "early_stopping": false,
145
+ "encoder_no_repeat_ngram_size": 0,
146
+ "eos_token_id": null,
147
+ "eval_capacity_factor": 1.4,
148
+ "exponential_decay_length_penalty": null,
149
+ "finetuning_task": null,
150
+ "forced_bos_token_id": null,
151
+ "forced_eos_token_id": null,
152
+ "hidden_act": "gelu",
153
+ "hidden_size": 1024,
154
+ "id2label": {
155
+ "0": "LABEL_0",
156
+ "1": "LABEL_1"
157
+ },
158
+ "image_size": 448,
159
+ "initializer_factor": 0.1,
160
+ "initializer_range": 1e-10,
161
+ "intermediate_size": 4096,
162
+ "is_decoder": false,
163
+ "is_encoder_decoder": false,
164
+ "label2id": {
165
+ "LABEL_0": 0,
166
+ "LABEL_1": 1
167
+ },
168
+ "laux_allreduce": "all_nodes",
169
+ "layer_norm_eps": 1e-06,
170
+ "length_penalty": 1.0,
171
+ "max_length": 20,
172
+ "min_length": 0,
173
+ "model_type": "intern_vit_6b",
174
+ "moe_coeff_ratio": 0.5,
175
+ "moe_intermediate_size": 768,
176
+ "moe_output_scale": 4.0,
177
+ "no_repeat_ngram_size": 0,
178
+ "noisy_gate_policy": "RSample_before",
179
+ "norm_type": "layer_norm",
180
+ "num_attention_heads": 16,
181
+ "num_beam_groups": 1,
182
+ "num_beams": 1,
183
+ "num_channels": 3,
184
+ "num_experts": 8,
185
+ "num_hidden_layers": 24,
186
+ "num_return_sequences": 1,
187
+ "num_routed_experts": 4,
188
+ "num_shared_experts": 4,
189
+ "output_attentions": false,
190
+ "output_hidden_states": false,
191
+ "output_scores": false,
192
+ "pad_token_id": null,
193
+ "patch_size": 14,
194
+ "prefix": null,
195
+ "problem_type": null,
196
+ "pruned_heads": {},
197
+ "qk_normalization": false,
198
+ "qkv_bias": true,
199
+ "remove_invalid_values": false,
200
+ "repetition_penalty": 1.0,
201
+ "return_dict": true,
202
+ "return_dict_in_generate": false,
203
+ "sep_token_id": null,
204
+ "shared_expert_intermediate_size": 3072,
205
+ "suppress_tokens": null,
206
+ "task_specific_params": null,
207
+ "temperature": 1.0,
208
+ "tf_legacy_loss": false,
209
+ "tie_encoder_decoder": false,
210
+ "tie_word_embeddings": true,
211
+ "tokenizer_class": null,
212
+ "top_k": 50,
213
+ "top_p": 1.0,
214
+ "torch_dtype": "bfloat16",
215
+ "torchscript": false,
216
+ "transformers_version": "4.48.3",
217
+ "typical_p": 1.0,
218
+ "use_bfloat16": false,
219
+ "use_flash_attn": true,
220
+ "use_moe": false,
221
+ "use_residual": true,
222
+ "use_rts": false,
223
+ "use_weighted_residual": false
224
+ }
225
+ }
native_backbone/configuration_intern_vit.py ADDED
@@ -0,0 +1,119 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # --------------------------------------------------------
2
+ # InternVL
3
+ # Copyright (c) 2024 OpenGVLab
4
+ # Licensed under The MIT License [see LICENSE for details]
5
+ # --------------------------------------------------------
6
+ import os
7
+ from typing import Union
8
+
9
+ from transformers.configuration_utils import PretrainedConfig
10
+ from transformers.utils import logging
11
+
12
+ logger = logging.get_logger(__name__)
13
+
14
+
15
+ class InternVisionConfig(PretrainedConfig):
16
+ r"""
17
+ This is the configuration class to store the configuration of a [`InternVisionModel`]. It is used to
18
+ instantiate a vision encoder according to the specified arguments, defining the model architecture.
19
+
20
+ Configuration objects inherit from [`PretrainedConfig`] and can be used to control the model outputs. Read the
21
+ documentation from [`PretrainedConfig`] for more information.
22
+
23
+ Args:
24
+ num_channels (`int`, *optional*, defaults to 3):
25
+ Number of color channels in the input images (e.g., 3 for RGB).
26
+ patch_size (`int`, *optional*, defaults to 14):
27
+ The size (resolution) of each patch.
28
+ image_size (`int`, *optional*, defaults to 224):
29
+ The size (resolution) of each image.
30
+ qkv_bias (`bool`, *optional*, defaults to `False`):
31
+ Whether to add a bias to the queries and values in the self-attention layers.
32
+ hidden_size (`int`, *optional*, defaults to 3200):
33
+ Dimensionality of the encoder layers and the pooler layer.
34
+ num_attention_heads (`int`, *optional*, defaults to 25):
35
+ Number of attention heads for each attention layer in the Transformer encoder.
36
+ intermediate_size (`int`, *optional*, defaults to 12800):
37
+ Dimensionality of the "intermediate" (i.e., feed-forward) layer in the Transformer encoder.
38
+ qk_normalization (`bool`, *optional*, defaults to `True`):
39
+ Whether to normalize the queries and keys in the self-attention layers.
40
+ num_hidden_layers (`int`, *optional*, defaults to 48):
41
+ Number of hidden layers in the Transformer encoder.
42
+ use_flash_attn (`bool`, *optional*, defaults to `True`):
43
+ Whether to use flash attention mechanism.
44
+ hidden_act (`str` or `function`, *optional*, defaults to `"gelu"`):
45
+ The non-linear activation function (function or string) in the encoder and pooler. If string, `"gelu"`,
46
+ `"relu"`, `"selu"` and `"gelu_new"` ``"gelu"` are supported.
47
+ layer_norm_eps (`float`, *optional*, defaults to 1e-6):
48
+ The epsilon used by the layer normalization layers.
49
+ dropout (`float`, *optional*, defaults to 0.0):
50
+ The dropout probability for all fully connected layers in the embeddings, encoder, and pooler.
51
+ drop_path_rate (`float`, *optional*, defaults to 0.0):
52
+ Dropout rate for stochastic depth.
53
+ attention_dropout (`float`, *optional*, defaults to 0.0):
54
+ The dropout ratio for the attention probabilities.
55
+ initializer_range (`float`, *optional*, defaults to 0.02):
56
+ The standard deviation of the truncated_normal_initializer for initializing all weight matrices.
57
+ initializer_factor (`float`, *optional*, defaults to 0.1):
58
+ A factor for layer scale.
59
+ """
60
+
61
+ model_type = 'intern_vit_6b'
62
+
63
+ def __init__(
64
+ self,
65
+ num_channels=3,
66
+ patch_size=14,
67
+ image_size=224,
68
+ qkv_bias=False,
69
+ hidden_size=3200,
70
+ num_attention_heads=25,
71
+ intermediate_size=12800,
72
+ qk_normalization=True,
73
+ num_hidden_layers=48,
74
+ use_flash_attn=True,
75
+ hidden_act='gelu',
76
+ norm_type='rms_norm',
77
+ layer_norm_eps=1e-6,
78
+ dropout=0.0,
79
+ drop_path_rate=0.0,
80
+ attention_dropout=0.0,
81
+ initializer_range=0.02,
82
+ initializer_factor=0.1,
83
+ **kwargs,
84
+ ):
85
+ super().__init__(**kwargs)
86
+
87
+ self.hidden_size = hidden_size
88
+ self.intermediate_size = intermediate_size
89
+ self.dropout = dropout
90
+ self.drop_path_rate = drop_path_rate
91
+ self.num_hidden_layers = num_hidden_layers
92
+ self.num_attention_heads = num_attention_heads
93
+ self.num_channels = num_channels
94
+ self.patch_size = patch_size
95
+ self.image_size = image_size
96
+ self.initializer_range = initializer_range
97
+ self.initializer_factor = initializer_factor
98
+ self.attention_dropout = attention_dropout
99
+ self.layer_norm_eps = layer_norm_eps
100
+ self.hidden_act = hidden_act
101
+ self.norm_type = norm_type
102
+ self.qkv_bias = qkv_bias
103
+ self.qk_normalization = qk_normalization
104
+ self.use_flash_attn = use_flash_attn
105
+
106
+ @classmethod
107
+ def from_pretrained(cls, pretrained_model_name_or_path: Union[str, os.PathLike], **kwargs) -> 'PretrainedConfig':
108
+ config_dict, kwargs = cls.get_config_dict(pretrained_model_name_or_path, **kwargs)
109
+
110
+ if 'vision_config' in config_dict:
111
+ config_dict = config_dict['vision_config']
112
+
113
+ if 'model_type' in config_dict and hasattr(cls, 'model_type') and config_dict['model_type'] != cls.model_type:
114
+ logger.warning(
115
+ f"You are using a model of type {config_dict['model_type']} to instantiate a model of type "
116
+ f'{cls.model_type}. This is not supported for all configurations of models and can yield errors.'
117
+ )
118
+
119
+ return cls.from_dict(config_dict, **kwargs)
native_backbone/configuration_internlm2.py ADDED
@@ -0,0 +1,150 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) The InternLM team and The HuggingFace Inc. team. All rights reserved.
2
+ #
3
+ # This code is based on transformers/src/transformers/models/llama/configuration_llama.py
4
+ #
5
+ # Licensed under the Apache License, Version 2.0 (the "License");
6
+ # you may not use this file except in compliance with the License.
7
+ # You may obtain a copy of the License at
8
+ #
9
+ # http://www.apache.org/licenses/LICENSE-2.0
10
+ #
11
+ # Unless required by applicable law or agreed to in writing, software
12
+ # distributed under the License is distributed on an "AS IS" BASIS,
13
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
14
+ # See the License for the specific language governing permissions and
15
+ # limitations under the License.
16
+ """ InternLM2 model configuration"""
17
+
18
+ from transformers.configuration_utils import PretrainedConfig
19
+ from transformers.utils import logging
20
+
21
+ logger = logging.get_logger(__name__)
22
+
23
+ INTERNLM2_PRETRAINED_CONFIG_ARCHIVE_MAP = {}
24
+
25
+
26
+ # Modified from transformers.model.llama.configuration_llama.LlamaConfig
27
+ class InternLM2Config(PretrainedConfig):
28
+ r"""
29
+ This is the configuration class to store the configuration of a [`InternLM2Model`]. It is used to instantiate
30
+ an InternLM2 model according to the specified arguments, defining the model architecture. Instantiating a
31
+ configuration with the defaults will yield a similar configuration to that of the InternLM2-7B.
32
+
33
+ Configuration objects inherit from [`PretrainedConfig`] and can be used to control the model outputs. Read the
34
+ documentation from [`PretrainedConfig`] for more information.
35
+
36
+
37
+ Args:
38
+ vocab_size (`int`, *optional*, defaults to 32000):
39
+ Vocabulary size of the InternLM2 model. Defines the number of different tokens that can be represented by the
40
+ `inputs_ids` passed when calling [`InternLM2Model`]
41
+ hidden_size (`int`, *optional*, defaults to 4096):
42
+ Dimension of the hidden representations.
43
+ intermediate_size (`int`, *optional*, defaults to 11008):
44
+ Dimension of the MLP representations.
45
+ num_hidden_layers (`int`, *optional*, defaults to 32):
46
+ Number of hidden layers in the Transformer encoder.
47
+ num_attention_heads (`int`, *optional*, defaults to 32):
48
+ Number of attention heads for each attention layer in the Transformer encoder.
49
+ num_key_value_heads (`int`, *optional*):
50
+ This is the number of key_value heads that should be used to implement Grouped Query Attention. If
51
+ `num_key_value_heads=num_attention_heads`, the model will use Multi Head Attention (MHA), if
52
+ `num_key_value_heads=1 the model will use Multi Query Attention (MQA) otherwise GQA is used. When
53
+ converting a multi-head checkpoint to a GQA checkpoint, each group key and value head should be constructed
54
+ by meanpooling all the original heads within that group. For more details checkout [this
55
+ paper](https://arxiv.org/pdf/2305.13245.pdf). If it is not specified, will default to
56
+ `num_attention_heads`.
57
+ hidden_act (`str` or `function`, *optional*, defaults to `"silu"`):
58
+ The non-linear activation function (function or string) in the decoder.
59
+ max_position_embeddings (`int`, *optional*, defaults to 2048):
60
+ The maximum sequence length that this model might ever be used with. Typically set this to something large
61
+ just in case (e.g., 512 or 1024 or 2048).
62
+ initializer_range (`float`, *optional*, defaults to 0.02):
63
+ The standard deviation of the truncated_normal_initializer for initializing all weight matrices.
64
+ rms_norm_eps (`float`, *optional*, defaults to 1e-12):
65
+ The epsilon used by the rms normalization layers.
66
+ use_cache (`bool`, *optional*, defaults to `True`):
67
+ Whether or not the model should return the last key/values attentions (not used by all models). Only
68
+ relevant if `config.is_decoder=True`.
69
+ tie_word_embeddings(`bool`, *optional*, defaults to `False`):
70
+ Whether to tie weight embeddings
71
+ Example:
72
+
73
+ """
74
+ model_type = 'internlm2'
75
+ _auto_class = 'AutoConfig'
76
+
77
+ def __init__( # pylint: disable=W0102
78
+ self,
79
+ vocab_size=103168,
80
+ hidden_size=4096,
81
+ intermediate_size=11008,
82
+ num_hidden_layers=32,
83
+ num_attention_heads=32,
84
+ num_key_value_heads=None,
85
+ hidden_act='silu',
86
+ max_position_embeddings=2048,
87
+ initializer_range=0.02,
88
+ rms_norm_eps=1e-6,
89
+ use_cache=True,
90
+ pad_token_id=0,
91
+ bos_token_id=1,
92
+ eos_token_id=2,
93
+ tie_word_embeddings=False,
94
+ bias=True,
95
+ rope_theta=10000,
96
+ rope_scaling=None,
97
+ attn_implementation='eager',
98
+ **kwargs,
99
+ ):
100
+ self.vocab_size = vocab_size
101
+ self.max_position_embeddings = max_position_embeddings
102
+ self.hidden_size = hidden_size
103
+ self.intermediate_size = intermediate_size
104
+ self.num_hidden_layers = num_hidden_layers
105
+ self.num_attention_heads = num_attention_heads
106
+ self.bias = bias
107
+
108
+ if num_key_value_heads is None:
109
+ num_key_value_heads = num_attention_heads
110
+ self.num_key_value_heads = num_key_value_heads
111
+
112
+ self.hidden_act = hidden_act
113
+ self.initializer_range = initializer_range
114
+ self.rms_norm_eps = rms_norm_eps
115
+ self.use_cache = use_cache
116
+ self.rope_theta = rope_theta
117
+ self.rope_scaling = rope_scaling
118
+ self._rope_scaling_validation()
119
+
120
+ self.attn_implementation = attn_implementation
121
+ if self.attn_implementation is None:
122
+ self.attn_implementation = 'eager'
123
+ super().__init__(
124
+ pad_token_id=pad_token_id,
125
+ bos_token_id=bos_token_id,
126
+ eos_token_id=eos_token_id,
127
+ tie_word_embeddings=tie_word_embeddings,
128
+ **kwargs,
129
+ )
130
+
131
+ def _rope_scaling_validation(self):
132
+ """
133
+ Validate the `rope_scaling` configuration.
134
+ """
135
+ if self.rope_scaling is None:
136
+ return
137
+
138
+ if not isinstance(self.rope_scaling, dict) or len(self.rope_scaling) != 2:
139
+ raise ValueError(
140
+ '`rope_scaling` must be a dictionary with with two fields, `type` and `factor`, '
141
+ f'got {self.rope_scaling}'
142
+ )
143
+ rope_scaling_type = self.rope_scaling.get('type', None)
144
+ rope_scaling_factor = self.rope_scaling.get('factor', None)
145
+ if rope_scaling_type is None or rope_scaling_type not in ['linear', 'dynamic']:
146
+ raise ValueError(
147
+ f"`rope_scaling`'s type field must be one of ['linear', 'dynamic'], got {rope_scaling_type}"
148
+ )
149
+ if rope_scaling_factor is None or not isinstance(rope_scaling_factor, float) or rope_scaling_factor < 1.0:
150
+ raise ValueError(f"`rope_scaling`'s factor field must be a float >= 1, got {rope_scaling_factor}")
native_backbone/configuration_internvl_chat.py ADDED
@@ -0,0 +1,101 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # --------------------------------------------------------
2
+ # InternVL
3
+ # Copyright (c) 2024 OpenGVLab
4
+ # Licensed under The MIT License [see LICENSE for details]
5
+ # --------------------------------------------------------
6
+
7
+ import copy
8
+
9
+ from transformers import AutoConfig, LlamaConfig, Qwen2Config
10
+ from transformers.configuration_utils import PretrainedConfig
11
+ from transformers.utils import logging
12
+
13
+ from .configuration_intern_vit import InternVisionConfig
14
+ from .configuration_internlm2 import InternLM2Config
15
+
16
+
17
+ logger = logging.get_logger(__name__)
18
+
19
+
20
+ class InternVLChatConfig(PretrainedConfig):
21
+ model_type = 'internvl_chat'
22
+ is_composition = True
23
+
24
+ def __init__(
25
+ self,
26
+ vision_config=None,
27
+ llm_config=None,
28
+ use_backbone_lora=0,
29
+ use_llm_lora=0,
30
+ select_layer=-1,
31
+ force_image_size=None,
32
+ downsample_ratio=0.5,
33
+ template=None,
34
+ dynamic_image_size=False,
35
+ use_thumbnail=False,
36
+ ps_version='v1',
37
+ min_dynamic_patch=1,
38
+ max_dynamic_patch=6,
39
+ **kwargs):
40
+ super().__init__(**kwargs)
41
+
42
+ if vision_config is None:
43
+ vision_config = {'architectures': ['InternVisionModel']}
44
+ logger.info('vision_config is None. Initializing the InternVisionConfig with default values.')
45
+
46
+ if llm_config is None:
47
+ llm_config = {'architectures': ['InternLM2ForCausalLM']}
48
+ logger.info('llm_config is None. Initializing the LlamaConfig config with default values (`LlamaConfig`).')
49
+
50
+ self.vision_config = InternVisionConfig(**vision_config)
51
+ if llm_config['architectures'][0] == 'LlamaForCausalLM':
52
+ self.llm_config = LlamaConfig(**llm_config)
53
+ elif llm_config['architectures'][0] == 'InternLM2ForCausalLM':
54
+ self.llm_config = InternLM2Config(**llm_config)
55
+ elif llm_config['architectures'][0] == 'Qwen2ForCausalLM':
56
+ self.llm_config = Qwen2Config(**llm_config)
57
+ else:
58
+ raise ValueError('Unsupported architecture: {}'.format(llm_config['architectures'][0]))
59
+ self.use_backbone_lora = use_backbone_lora
60
+ self.use_llm_lora = use_llm_lora
61
+ self.select_layer = select_layer
62
+ self.force_image_size = force_image_size
63
+ self.downsample_ratio = downsample_ratio
64
+ self.template = template
65
+ self.dynamic_image_size = dynamic_image_size
66
+ self.use_thumbnail = use_thumbnail
67
+ self.ps_version = ps_version # pixel shuffle version
68
+ self.min_dynamic_patch = min_dynamic_patch
69
+ self.max_dynamic_patch = max_dynamic_patch
70
+ # By default, we use tie_word_embeddings=False for models of all sizes.
71
+ self.tie_word_embeddings = self.llm_config.tie_word_embeddings
72
+
73
+ logger.info(f'vision_select_layer: {self.select_layer}')
74
+ logger.info(f'ps_version: {self.ps_version}')
75
+ logger.info(f'min_dynamic_patch: {self.min_dynamic_patch}')
76
+ logger.info(f'max_dynamic_patch: {self.max_dynamic_patch}')
77
+
78
+ def to_dict(self):
79
+ """
80
+ Serializes this instance to a Python dictionary. Override the default [`~PretrainedConfig.to_dict`].
81
+
82
+ Returns:
83
+ `Dict[str, any]`: Dictionary of all the attributes that make up this configuration instance,
84
+ """
85
+ output = copy.deepcopy(self.__dict__)
86
+ output['vision_config'] = self.vision_config.to_dict()
87
+ output['llm_config'] = self.llm_config.to_dict()
88
+ output['model_type'] = self.__class__.model_type
89
+ output['use_backbone_lora'] = self.use_backbone_lora
90
+ output['use_llm_lora'] = self.use_llm_lora
91
+ output['select_layer'] = self.select_layer
92
+ output['force_image_size'] = self.force_image_size
93
+ output['downsample_ratio'] = self.downsample_ratio
94
+ output['template'] = self.template
95
+ output['dynamic_image_size'] = self.dynamic_image_size
96
+ output['use_thumbnail'] = self.use_thumbnail
97
+ output['ps_version'] = self.ps_version
98
+ output['min_dynamic_patch'] = self.min_dynamic_patch
99
+ output['max_dynamic_patch'] = self.max_dynamic_patch
100
+
101
+ return output
native_backbone/conversation.py ADDED
@@ -0,0 +1,391 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ Conversation prompt templates.
3
+
4
+ We kindly request that you import fastchat instead of copying this file if you wish to use it.
5
+ If you have changes in mind, please contribute back so the community can benefit collectively and continue to maintain these valuable templates.
6
+
7
+ Modified from https://github.com/lm-sys/FastChat/blob/main/fastchat/conversation.py
8
+ """
9
+
10
+ import dataclasses
11
+ from enum import IntEnum, auto
12
+ from typing import Dict, List, Tuple, Union
13
+
14
+
15
+ class SeparatorStyle(IntEnum):
16
+ """Separator styles."""
17
+
18
+ ADD_COLON_SINGLE = auto()
19
+ ADD_COLON_TWO = auto()
20
+ ADD_COLON_SPACE_SINGLE = auto()
21
+ NO_COLON_SINGLE = auto()
22
+ NO_COLON_TWO = auto()
23
+ ADD_NEW_LINE_SINGLE = auto()
24
+ LLAMA2 = auto()
25
+ CHATGLM = auto()
26
+ CHATML = auto()
27
+ CHATINTERN = auto()
28
+ DOLLY = auto()
29
+ RWKV = auto()
30
+ PHOENIX = auto()
31
+ ROBIN = auto()
32
+ FALCON_CHAT = auto()
33
+ CHATGLM3 = auto()
34
+ INTERNVL_ZH = auto()
35
+ MPT = auto()
36
+
37
+
38
+ @dataclasses.dataclass
39
+ class Conversation:
40
+ """A class that manages prompt templates and keeps all conversation history."""
41
+
42
+ # The name of this template
43
+ name: str
44
+ # The template of the system prompt
45
+ system_template: str = '{system_message}'
46
+ # The system message
47
+ system_message: str = ''
48
+ # The names of two roles
49
+ roles: Tuple[str] = ('USER', 'ASSISTANT')
50
+ # All messages. Each item is (role, message).
51
+ messages: List[List[str]] = ()
52
+ # The number of few shot examples
53
+ offset: int = 0
54
+ # The separator style and configurations
55
+ sep_style: SeparatorStyle = SeparatorStyle.ADD_COLON_SINGLE
56
+ sep: str = '\n'
57
+ sep2: str = None
58
+ # Stop criteria (the default one is EOS token)
59
+ stop_str: Union[str, List[str]] = None
60
+ # Stops generation if meeting any token in this list
61
+ stop_token_ids: List[int] = None
62
+
63
+ def get_prompt(self) -> str:
64
+ """Get the prompt for generation."""
65
+ system_prompt = self.system_template.format(system_message=self.system_message)
66
+ if self.sep_style == SeparatorStyle.ADD_COLON_SINGLE:
67
+ ret = system_prompt + self.sep
68
+ for role, message in self.messages:
69
+ if message:
70
+ ret += role + ': ' + message + self.sep
71
+ else:
72
+ ret += role + ':'
73
+ return ret
74
+ elif self.sep_style == SeparatorStyle.ADD_COLON_TWO:
75
+ seps = [self.sep, self.sep2]
76
+ ret = system_prompt + seps[0]
77
+ for i, (role, message) in enumerate(self.messages):
78
+ if message:
79
+ ret += role + ': ' + message + seps[i % 2]
80
+ else:
81
+ ret += role + ':'
82
+ return ret
83
+ elif self.sep_style == SeparatorStyle.ADD_COLON_SPACE_SINGLE:
84
+ ret = system_prompt + self.sep
85
+ for role, message in self.messages:
86
+ if message:
87
+ ret += role + ': ' + message + self.sep
88
+ else:
89
+ ret += role + ': ' # must be end with a space
90
+ return ret
91
+ elif self.sep_style == SeparatorStyle.ADD_NEW_LINE_SINGLE:
92
+ ret = '' if system_prompt == '' else system_prompt + self.sep
93
+ for role, message in self.messages:
94
+ if message:
95
+ ret += role + '\n' + message + self.sep
96
+ else:
97
+ ret += role + '\n'
98
+ return ret
99
+ elif self.sep_style == SeparatorStyle.NO_COLON_SINGLE:
100
+ ret = system_prompt
101
+ for role, message in self.messages:
102
+ if message:
103
+ ret += role + message + self.sep
104
+ else:
105
+ ret += role
106
+ return ret
107
+ elif self.sep_style == SeparatorStyle.NO_COLON_TWO:
108
+ seps = [self.sep, self.sep2]
109
+ ret = system_prompt
110
+ for i, (role, message) in enumerate(self.messages):
111
+ if message:
112
+ ret += role + message + seps[i % 2]
113
+ else:
114
+ ret += role
115
+ return ret
116
+ elif self.sep_style == SeparatorStyle.RWKV:
117
+ ret = system_prompt
118
+ for i, (role, message) in enumerate(self.messages):
119
+ if message:
120
+ ret += (
121
+ role
122
+ + ': '
123
+ + message.replace('\r\n', '\n').replace('\n\n', '\n')
124
+ )
125
+ ret += '\n\n'
126
+ else:
127
+ ret += role + ':'
128
+ return ret
129
+ elif self.sep_style == SeparatorStyle.LLAMA2:
130
+ seps = [self.sep, self.sep2]
131
+ if self.system_message:
132
+ ret = system_prompt
133
+ else:
134
+ ret = '[INST] '
135
+ for i, (role, message) in enumerate(self.messages):
136
+ tag = self.roles[i % 2]
137
+ if message:
138
+ if i == 0:
139
+ ret += message + ' '
140
+ else:
141
+ ret += tag + ' ' + message + seps[i % 2]
142
+ else:
143
+ ret += tag
144
+ return ret
145
+ elif self.sep_style == SeparatorStyle.CHATGLM:
146
+ # source: https://huggingface.co/THUDM/chatglm-6b/blob/1d240ba371910e9282298d4592532d7f0f3e9f3e/modeling_chatglm.py#L1302-L1308
147
+ # source2: https://huggingface.co/THUDM/chatglm2-6b/blob/e186c891cf64310ac66ef10a87e6635fa6c2a579/modeling_chatglm.py#L926
148
+ round_add_n = 1 if self.name == 'chatglm2' else 0
149
+ if system_prompt:
150
+ ret = system_prompt + self.sep
151
+ else:
152
+ ret = ''
153
+
154
+ for i, (role, message) in enumerate(self.messages):
155
+ if i % 2 == 0:
156
+ ret += f'[Round {i//2 + round_add_n}]{self.sep}'
157
+
158
+ if message:
159
+ ret += f'{role}:{message}{self.sep}'
160
+ else:
161
+ ret += f'{role}:'
162
+ return ret
163
+ elif self.sep_style == SeparatorStyle.CHATML:
164
+ ret = '' if system_prompt == '' else system_prompt + self.sep + '\n'
165
+ for role, message in self.messages:
166
+ if message:
167
+ ret += role + '\n' + message + self.sep + '\n'
168
+ else:
169
+ ret += role + '\n'
170
+ return ret
171
+ elif self.sep_style == SeparatorStyle.CHATGLM3:
172
+ ret = ''
173
+ if self.system_message:
174
+ ret += system_prompt
175
+ for role, message in self.messages:
176
+ if message:
177
+ ret += role + '\n' + ' ' + message
178
+ else:
179
+ ret += role
180
+ return ret
181
+ elif self.sep_style == SeparatorStyle.CHATINTERN:
182
+ # source: https://huggingface.co/internlm/internlm-chat-7b-8k/blob/bd546fa984b4b0b86958f56bf37f94aa75ab8831/modeling_internlm.py#L771
183
+ seps = [self.sep, self.sep2]
184
+ ret = system_prompt
185
+ for i, (role, message) in enumerate(self.messages):
186
+ # if i % 2 == 0:
187
+ # ret += "<s>"
188
+ if message:
189
+ ret += role + ':' + message + seps[i % 2] + '\n'
190
+ else:
191
+ ret += role + ':'
192
+ return ret
193
+ elif self.sep_style == SeparatorStyle.DOLLY:
194
+ seps = [self.sep, self.sep2]
195
+ ret = system_prompt
196
+ for i, (role, message) in enumerate(self.messages):
197
+ if message:
198
+ ret += role + ':\n' + message + seps[i % 2]
199
+ if i % 2 == 1:
200
+ ret += '\n\n'
201
+ else:
202
+ ret += role + ':\n'
203
+ return ret
204
+ elif self.sep_style == SeparatorStyle.PHOENIX:
205
+ ret = system_prompt
206
+ for role, message in self.messages:
207
+ if message:
208
+ ret += role + ': ' + '<s>' + message + '</s>'
209
+ else:
210
+ ret += role + ': ' + '<s>'
211
+ return ret
212
+ elif self.sep_style == SeparatorStyle.ROBIN:
213
+ ret = system_prompt + self.sep
214
+ for role, message in self.messages:
215
+ if message:
216
+ ret += role + ':\n' + message + self.sep
217
+ else:
218
+ ret += role + ':\n'
219
+ return ret
220
+ elif self.sep_style == SeparatorStyle.FALCON_CHAT:
221
+ ret = ''
222
+ if self.system_message:
223
+ ret += system_prompt + self.sep
224
+ for role, message in self.messages:
225
+ if message:
226
+ ret += role + ': ' + message + self.sep
227
+ else:
228
+ ret += role + ':'
229
+
230
+ return ret
231
+ elif self.sep_style == SeparatorStyle.INTERNVL_ZH:
232
+ seps = [self.sep, self.sep2]
233
+ ret = self.system_message + seps[0]
234
+ for i, (role, message) in enumerate(self.messages):
235
+ if message:
236
+ ret += role + ': ' + message + seps[i % 2]
237
+ else:
238
+ ret += role + ':'
239
+ return ret
240
+ elif self.sep_style == SeparatorStyle.MPT:
241
+ ret = system_prompt + self.sep
242
+ for role, message in self.messages:
243
+ if message:
244
+ if type(message) is tuple:
245
+ message, _, _ = message
246
+ ret += role + message + self.sep
247
+ else:
248
+ ret += role
249
+ return ret
250
+ else:
251
+ raise ValueError(f'Invalid style: {self.sep_style}')
252
+
253
+ def set_system_message(self, system_message: str):
254
+ """Set the system message."""
255
+ self.system_message = system_message
256
+
257
+ def append_message(self, role: str, message: str):
258
+ """Append a new message."""
259
+ self.messages.append([role, message])
260
+
261
+ def update_last_message(self, message: str):
262
+ """Update the last output.
263
+
264
+ The last message is typically set to be None when constructing the prompt,
265
+ so we need to update it in-place after getting the response from a model.
266
+ """
267
+ self.messages[-1][1] = message
268
+
269
+ def to_gradio_chatbot(self):
270
+ """Convert the conversation to gradio chatbot format."""
271
+ ret = []
272
+ for i, (role, msg) in enumerate(self.messages[self.offset :]):
273
+ if i % 2 == 0:
274
+ ret.append([msg, None])
275
+ else:
276
+ ret[-1][-1] = msg
277
+ return ret
278
+
279
+ def to_openai_api_messages(self):
280
+ """Convert the conversation to OpenAI chat completion format."""
281
+ ret = [{'role': 'system', 'content': self.system_message}]
282
+
283
+ for i, (_, msg) in enumerate(self.messages[self.offset :]):
284
+ if i % 2 == 0:
285
+ ret.append({'role': 'user', 'content': msg})
286
+ else:
287
+ if msg is not None:
288
+ ret.append({'role': 'assistant', 'content': msg})
289
+ return ret
290
+
291
+ def copy(self):
292
+ return Conversation(
293
+ name=self.name,
294
+ system_template=self.system_template,
295
+ system_message=self.system_message,
296
+ roles=self.roles,
297
+ messages=[[x, y] for x, y in self.messages],
298
+ offset=self.offset,
299
+ sep_style=self.sep_style,
300
+ sep=self.sep,
301
+ sep2=self.sep2,
302
+ stop_str=self.stop_str,
303
+ stop_token_ids=self.stop_token_ids,
304
+ )
305
+
306
+ def dict(self):
307
+ return {
308
+ 'template_name': self.name,
309
+ 'system_message': self.system_message,
310
+ 'roles': self.roles,
311
+ 'messages': self.messages,
312
+ 'offset': self.offset,
313
+ }
314
+
315
+
316
+ # A global registry for all conversation templates
317
+ conv_templates: Dict[str, Conversation] = {}
318
+
319
+
320
+ def register_conv_template(template: Conversation, override: bool = False):
321
+ """Register a new conversation template."""
322
+ if not override:
323
+ assert (
324
+ template.name not in conv_templates
325
+ ), f'{template.name} has been registered.'
326
+
327
+ conv_templates[template.name] = template
328
+
329
+
330
+ def get_conv_template(name: str) -> Conversation:
331
+ """Get a conversation template."""
332
+ return conv_templates[name].copy()
333
+
334
+
335
+ # Both Hermes-2 and internlm2-chat are chatml-format conversation templates. The difference
336
+ # is that during training, the preprocessing function for the Hermes-2 template doesn't add
337
+ # <s> at the beginning of the tokenized sequence, while the internlm2-chat template does.
338
+ # Therefore, they are completely equivalent during inference.
339
+ register_conv_template(
340
+ Conversation(
341
+ name='Hermes-2',
342
+ system_template='<|im_start|>system\n{system_message}',
343
+ # note: The new system prompt was not used here to avoid changes in benchmark performance.
344
+ # system_message='我是书生·万象,英文名是InternVL,是由上海人工智能实验室、清华大学及多家合作单位联合开发的多模态大语言模型。',
345
+ system_message='你是由上海人工智能实验室联合商汤科技开发的书生多模态大模型,英文名叫InternVL, 是一个有用无害的人工智能助手。',
346
+ roles=('<|im_start|>user\n', '<|im_start|>assistant\n'),
347
+ sep_style=SeparatorStyle.MPT,
348
+ sep='<|im_end|>',
349
+ stop_str='<|endoftext|>',
350
+ )
351
+ )
352
+
353
+
354
+ register_conv_template(
355
+ Conversation(
356
+ name='internlm2-chat',
357
+ system_template='<|im_start|>system\n{system_message}',
358
+ # note: The new system prompt was not used here to avoid changes in benchmark performance.
359
+ # system_message='我是书生·万象,英文名是InternVL,是由上海人工智能实验室、清华大学及多家合作单位联合开发的多模态大语言模型。',
360
+ system_message='你是由上海人工智能实验室联合商汤科技开发的书生多模态大模型,英文名叫InternVL, 是一个有用无害的人工智能助手。',
361
+ roles=('<|im_start|>user\n', '<|im_start|>assistant\n'),
362
+ sep_style=SeparatorStyle.MPT,
363
+ sep='<|im_end|>',
364
+ )
365
+ )
366
+
367
+
368
+ register_conv_template(
369
+ Conversation(
370
+ name='phi3-chat',
371
+ system_template='<|system|>\n{system_message}',
372
+ # note: The new system prompt was not used here to avoid changes in benchmark performance.
373
+ # system_message='我是书生·万象,英文名是InternVL,是由上海人工智能实验室、清华大学及多家合作单位联合开发的多模态大语言模型。',
374
+ system_message='你是由上海人工智能实验室联合商汤科技开发的书生多模态大模型,英文名叫InternVL, 是一个有用无害的人工智能助手。',
375
+ roles=('<|user|>\n', '<|assistant|>\n'),
376
+ sep_style=SeparatorStyle.MPT,
377
+ sep='<|end|>',
378
+ )
379
+ )
380
+
381
+
382
+ register_conv_template(
383
+ Conversation(
384
+ name='internvl2_5',
385
+ system_template='<|im_start|>system\n{system_message}',
386
+ system_message='你是书生·万象,英文名是InternVL,是由上海人工智能实验室、清华大学及多家合作单位联合开发的多模态大语言模型。',
387
+ roles=('<|im_start|>user\n', '<|im_start|>assistant\n'),
388
+ sep_style=SeparatorStyle.MPT,
389
+ sep='<|im_end|>\n',
390
+ )
391
+ )
native_backbone/generation_config.json ADDED
@@ -0,0 +1,4 @@
 
 
 
 
 
1
+ {
2
+ "_from_model_config": true,
3
+ "transformers_version": "4.48.3"
4
+ }
native_backbone/model.safetensors.index.json ADDED
@@ -0,0 +1,692 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "metadata": {
3
+ "total_size": 18277586944
4
+ },
5
+ "weight_map": {
6
+ "language_model.model.layers.0.attention.wo.weight": "native-00001.safetensors",
7
+ "language_model.model.layers.0.attention.wqkv.weight": "native-00001.safetensors",
8
+ "language_model.model.layers.0.attention_norm.weight": "native-00001.safetensors",
9
+ "language_model.model.layers.0.feed_forward.w1.weight": "native-00001.safetensors",
10
+ "language_model.model.layers.0.feed_forward.w2.weight": "native-00001.safetensors",
11
+ "language_model.model.layers.0.feed_forward.w3.weight": "native-00001.safetensors",
12
+ "language_model.model.layers.0.ffn_norm.weight": "native-00001.safetensors",
13
+ "language_model.model.layers.1.attention.wo.weight": "native-00001.safetensors",
14
+ "language_model.model.layers.1.attention.wqkv.weight": "native-00001.safetensors",
15
+ "language_model.model.layers.1.attention_norm.weight": "native-00001.safetensors",
16
+ "language_model.model.layers.1.feed_forward.w1.weight": "native-00001.safetensors",
17
+ "language_model.model.layers.1.feed_forward.w2.weight": "native-00001.safetensors",
18
+ "language_model.model.layers.1.feed_forward.w3.weight": "native-00001.safetensors",
19
+ "language_model.model.layers.1.ffn_norm.weight": "native-00001.safetensors",
20
+ "language_model.model.layers.10.attention.wo.weight": "native-00001.safetensors",
21
+ "language_model.model.layers.10.attention.wqkv.weight": "native-00001.safetensors",
22
+ "language_model.model.layers.10.attention_norm.weight": "native-00001.safetensors",
23
+ "language_model.model.layers.10.feed_forward.w1.weight": "native-00001.safetensors",
24
+ "language_model.model.layers.10.feed_forward.w2.weight": "native-00001.safetensors",
25
+ "language_model.model.layers.10.feed_forward.w3.weight": "native-00001.safetensors",
26
+ "language_model.model.layers.10.ffn_norm.weight": "native-00001.safetensors",
27
+ "language_model.model.layers.11.attention.wo.weight": "native-00001.safetensors",
28
+ "language_model.model.layers.11.attention.wqkv.weight": "native-00001.safetensors",
29
+ "language_model.model.layers.11.attention_norm.weight": "native-00001.safetensors",
30
+ "language_model.model.layers.11.feed_forward.w1.weight": "native-00001.safetensors",
31
+ "language_model.model.layers.11.feed_forward.w2.weight": "native-00001.safetensors",
32
+ "language_model.model.layers.11.feed_forward.w3.weight": "native-00001.safetensors",
33
+ "language_model.model.layers.11.ffn_norm.weight": "native-00001.safetensors",
34
+ "language_model.model.layers.12.attention.wo.weight": "native-00001.safetensors",
35
+ "language_model.model.layers.12.attention.wqkv.weight": "native-00001.safetensors",
36
+ "language_model.model.layers.12.attention_norm.weight": "native-00001.safetensors",
37
+ "language_model.model.layers.12.feed_forward.w1.weight": "native-00001.safetensors",
38
+ "language_model.model.layers.12.feed_forward.w2.weight": "native-00001.safetensors",
39
+ "language_model.model.layers.12.feed_forward.w3.weight": "native-00001.safetensors",
40
+ "language_model.model.layers.12.ffn_norm.weight": "native-00001.safetensors",
41
+ "language_model.model.layers.13.attention.wo.weight": "native-00001.safetensors",
42
+ "language_model.model.layers.13.attention.wqkv.weight": "native-00001.safetensors",
43
+ "language_model.model.layers.13.attention_norm.weight": "native-00001.safetensors",
44
+ "language_model.model.layers.13.feed_forward.w1.weight": "native-00001.safetensors",
45
+ "language_model.model.layers.13.feed_forward.w2.weight": "native-00001.safetensors",
46
+ "language_model.model.layers.13.feed_forward.w3.weight": "native-00001.safetensors",
47
+ "language_model.model.layers.13.ffn_norm.weight": "native-00001.safetensors",
48
+ "language_model.model.layers.14.attention.wo.weight": "native-00001.safetensors",
49
+ "language_model.model.layers.14.attention.wqkv.weight": "native-00001.safetensors",
50
+ "language_model.model.layers.14.attention_norm.weight": "native-00001.safetensors",
51
+ "language_model.model.layers.14.feed_forward.w1.weight": "native-00001.safetensors",
52
+ "language_model.model.layers.14.feed_forward.w2.weight": "native-00001.safetensors",
53
+ "language_model.model.layers.14.feed_forward.w3.weight": "native-00001.safetensors",
54
+ "language_model.model.layers.14.ffn_norm.weight": "native-00001.safetensors",
55
+ "language_model.model.layers.15.attention.wo.weight": "native-00001.safetensors",
56
+ "language_model.model.layers.15.attention.wqkv.weight": "native-00001.safetensors",
57
+ "language_model.model.layers.15.attention_norm.weight": "native-00001.safetensors",
58
+ "language_model.model.layers.15.feed_forward.w1.weight": "native-00001.safetensors",
59
+ "language_model.model.layers.15.feed_forward.w2.weight": "native-00001.safetensors",
60
+ "language_model.model.layers.15.feed_forward.w3.weight": "native-00001.safetensors",
61
+ "language_model.model.layers.15.ffn_norm.weight": "native-00001.safetensors",
62
+ "language_model.model.layers.16.attention.wo.weight": "native-00001.safetensors",
63
+ "language_model.model.layers.16.attention.wqkv.weight": "native-00001.safetensors",
64
+ "language_model.model.layers.16.attention_norm.weight": "native-00001.safetensors",
65
+ "language_model.model.layers.16.feed_forward.w1.weight": "native-00001.safetensors",
66
+ "language_model.model.layers.16.feed_forward.w2.weight": "native-00001.safetensors",
67
+ "language_model.model.layers.16.feed_forward.w3.weight": "native-00001.safetensors",
68
+ "language_model.model.layers.16.ffn_norm.weight": "native-00001.safetensors",
69
+ "language_model.model.layers.17.attention.wo.weight": "native-00001.safetensors",
70
+ "language_model.model.layers.17.attention.wqkv.weight": "native-00001.safetensors",
71
+ "language_model.model.layers.17.attention_norm.weight": "native-00001.safetensors",
72
+ "language_model.model.layers.17.feed_forward.w1.weight": "native-00001.safetensors",
73
+ "language_model.model.layers.17.feed_forward.w2.weight": "native-00001.safetensors",
74
+ "language_model.model.layers.17.feed_forward.w3.weight": "native-00001.safetensors",
75
+ "language_model.model.layers.17.ffn_norm.weight": "native-00001.safetensors",
76
+ "language_model.model.layers.18.attention.wo.weight": "native-00001.safetensors",
77
+ "language_model.model.layers.18.attention.wqkv.weight": "native-00001.safetensors",
78
+ "language_model.model.layers.18.attention_norm.weight": "native-00001.safetensors",
79
+ "language_model.model.layers.18.feed_forward.w1.weight": "native-00001.safetensors",
80
+ "language_model.model.layers.18.feed_forward.w2.weight": "native-00001.safetensors",
81
+ "language_model.model.layers.18.feed_forward.w3.weight": "native-00001.safetensors",
82
+ "language_model.model.layers.18.ffn_norm.weight": "native-00001.safetensors",
83
+ "language_model.model.layers.19.attention.wo.weight": "native-00001.safetensors",
84
+ "language_model.model.layers.19.attention.wqkv.weight": "native-00001.safetensors",
85
+ "language_model.model.layers.19.attention_norm.weight": "native-00001.safetensors",
86
+ "language_model.model.layers.19.feed_forward.w1.weight": "native-00001.safetensors",
87
+ "language_model.model.layers.19.feed_forward.w2.weight": "native-00001.safetensors",
88
+ "language_model.model.layers.19.feed_forward.w3.weight": "native-00001.safetensors",
89
+ "language_model.model.layers.19.ffn_norm.weight": "native-00001.safetensors",
90
+ "language_model.model.layers.2.attention.wo.weight": "native-00001.safetensors",
91
+ "language_model.model.layers.2.attention.wqkv.weight": "native-00001.safetensors",
92
+ "language_model.model.layers.2.attention_norm.weight": "native-00001.safetensors",
93
+ "language_model.model.layers.2.feed_forward.w1.weight": "native-00001.safetensors",
94
+ "language_model.model.layers.2.feed_forward.w2.weight": "native-00001.safetensors",
95
+ "language_model.model.layers.2.feed_forward.w3.weight": "native-00001.safetensors",
96
+ "language_model.model.layers.2.ffn_norm.weight": "native-00001.safetensors",
97
+ "language_model.model.layers.20.attention.wo.weight": "native-00001.safetensors",
98
+ "language_model.model.layers.20.attention.wqkv.weight": "native-00001.safetensors",
99
+ "language_model.model.layers.20.attention_norm.weight": "native-00001.safetensors",
100
+ "language_model.model.layers.20.feed_forward.w1.weight": "native-00001.safetensors",
101
+ "language_model.model.layers.20.feed_forward.w2.weight": "native-00001.safetensors",
102
+ "language_model.model.layers.20.feed_forward.w3.weight": "native-00001.safetensors",
103
+ "language_model.model.layers.20.ffn_norm.weight": "native-00001.safetensors",
104
+ "language_model.model.layers.21.attention.wo.weight": "native-00001.safetensors",
105
+ "language_model.model.layers.21.attention.wqkv.weight": "native-00001.safetensors",
106
+ "language_model.model.layers.21.attention_norm.weight": "native-00001.safetensors",
107
+ "language_model.model.layers.21.feed_forward.w1.weight": "native-00001.safetensors",
108
+ "language_model.model.layers.21.feed_forward.w2.weight": "native-00001.safetensors",
109
+ "language_model.model.layers.21.feed_forward.w3.weight": "native-00001.safetensors",
110
+ "language_model.model.layers.21.ffn_norm.weight": "native-00001.safetensors",
111
+ "language_model.model.layers.22.attention.wo.weight": "native-00001.safetensors",
112
+ "language_model.model.layers.22.attention.wqkv.weight": "native-00002.safetensors",
113
+ "language_model.model.layers.22.attention_norm.weight": "native-00002.safetensors",
114
+ "language_model.model.layers.22.feed_forward.w1.weight": "native-00002.safetensors",
115
+ "language_model.model.layers.22.feed_forward.w2.weight": "native-00002.safetensors",
116
+ "language_model.model.layers.22.feed_forward.w3.weight": "native-00002.safetensors",
117
+ "language_model.model.layers.22.ffn_norm.weight": "native-00002.safetensors",
118
+ "language_model.model.layers.23.attention.wo.weight": "native-00002.safetensors",
119
+ "language_model.model.layers.23.attention.wqkv.weight": "native-00002.safetensors",
120
+ "language_model.model.layers.23.attention_norm.weight": "native-00002.safetensors",
121
+ "language_model.model.layers.23.feed_forward.w1.weight": "native-00002.safetensors",
122
+ "language_model.model.layers.23.feed_forward.w2.weight": "native-00002.safetensors",
123
+ "language_model.model.layers.23.feed_forward.w3.weight": "native-00002.safetensors",
124
+ "language_model.model.layers.23.ffn_norm.weight": "native-00002.safetensors",
125
+ "language_model.model.layers.24.attention.wo.weight": "native-00002.safetensors",
126
+ "language_model.model.layers.24.attention.wqkv.weight": "native-00002.safetensors",
127
+ "language_model.model.layers.24.attention_norm.weight": "native-00002.safetensors",
128
+ "language_model.model.layers.24.feed_forward.w1.weight": "native-00002.safetensors",
129
+ "language_model.model.layers.24.feed_forward.w2.weight": "native-00002.safetensors",
130
+ "language_model.model.layers.24.feed_forward.w3.weight": "native-00002.safetensors",
131
+ "language_model.model.layers.24.ffn_norm.weight": "native-00002.safetensors",
132
+ "language_model.model.layers.25.attention.wo.weight": "native-00002.safetensors",
133
+ "language_model.model.layers.25.attention.wqkv.weight": "native-00002.safetensors",
134
+ "language_model.model.layers.25.attention_norm.weight": "native-00002.safetensors",
135
+ "language_model.model.layers.25.feed_forward.w1.weight": "native-00002.safetensors",
136
+ "language_model.model.layers.25.feed_forward.w2.weight": "native-00002.safetensors",
137
+ "language_model.model.layers.25.feed_forward.w3.weight": "native-00002.safetensors",
138
+ "language_model.model.layers.25.ffn_norm.weight": "native-00002.safetensors",
139
+ "language_model.model.layers.26.attention.wo.weight": "native-00002.safetensors",
140
+ "language_model.model.layers.26.attention.wqkv.weight": "native-00002.safetensors",
141
+ "language_model.model.layers.26.attention_norm.weight": "native-00002.safetensors",
142
+ "language_model.model.layers.26.feed_forward.w1.weight": "native-00002.safetensors",
143
+ "language_model.model.layers.26.feed_forward.w2.weight": "native-00002.safetensors",
144
+ "language_model.model.layers.26.feed_forward.w3.weight": "native-00002.safetensors",
145
+ "language_model.model.layers.26.ffn_norm.weight": "native-00002.safetensors",
146
+ "language_model.model.layers.27.attention.wo.weight": "native-00002.safetensors",
147
+ "language_model.model.layers.27.attention.wqkv.weight": "native-00002.safetensors",
148
+ "language_model.model.layers.27.attention_norm.weight": "native-00002.safetensors",
149
+ "language_model.model.layers.27.feed_forward.w1.weight": "native-00002.safetensors",
150
+ "language_model.model.layers.27.feed_forward.w2.weight": "native-00002.safetensors",
151
+ "language_model.model.layers.27.feed_forward.w3.weight": "native-00002.safetensors",
152
+ "language_model.model.layers.27.ffn_norm.weight": "native-00002.safetensors",
153
+ "language_model.model.layers.28.attention.wo.weight": "native-00002.safetensors",
154
+ "language_model.model.layers.28.attention.wqkv.weight": "native-00002.safetensors",
155
+ "language_model.model.layers.28.attention_norm.weight": "native-00002.safetensors",
156
+ "language_model.model.layers.28.feed_forward.w1.weight": "native-00002.safetensors",
157
+ "language_model.model.layers.28.feed_forward.w2.weight": "native-00002.safetensors",
158
+ "language_model.model.layers.28.feed_forward.w3.weight": "native-00002.safetensors",
159
+ "language_model.model.layers.28.ffn_norm.weight": "native-00002.safetensors",
160
+ "language_model.model.layers.29.attention.wo.weight": "native-00002.safetensors",
161
+ "language_model.model.layers.29.attention.wqkv.weight": "native-00002.safetensors",
162
+ "language_model.model.layers.29.attention_norm.weight": "native-00002.safetensors",
163
+ "language_model.model.layers.29.feed_forward.w1.weight": "native-00002.safetensors",
164
+ "language_model.model.layers.29.feed_forward.w2.weight": "native-00002.safetensors",
165
+ "language_model.model.layers.29.feed_forward.w3.weight": "native-00002.safetensors",
166
+ "language_model.model.layers.29.ffn_norm.weight": "native-00002.safetensors",
167
+ "language_model.model.layers.3.attention.wo.weight": "native-00002.safetensors",
168
+ "language_model.model.layers.3.attention.wqkv.weight": "native-00002.safetensors",
169
+ "language_model.model.layers.3.attention_norm.weight": "native-00002.safetensors",
170
+ "language_model.model.layers.3.feed_forward.w1.weight": "native-00002.safetensors",
171
+ "language_model.model.layers.3.feed_forward.w2.weight": "native-00002.safetensors",
172
+ "language_model.model.layers.3.feed_forward.w3.weight": "native-00002.safetensors",
173
+ "language_model.model.layers.3.ffn_norm.weight": "native-00002.safetensors",
174
+ "language_model.model.layers.30.attention.wo.weight": "native-00002.safetensors",
175
+ "language_model.model.layers.30.attention.wqkv.weight": "native-00002.safetensors",
176
+ "language_model.model.layers.30.attention_norm.weight": "native-00002.safetensors",
177
+ "language_model.model.layers.30.feed_forward.w1.weight": "native-00002.safetensors",
178
+ "language_model.model.layers.30.feed_forward.w2.weight": "native-00002.safetensors",
179
+ "language_model.model.layers.30.feed_forward.w3.weight": "native-00002.safetensors",
180
+ "language_model.model.layers.30.ffn_norm.weight": "native-00002.safetensors",
181
+ "language_model.model.layers.31.attention.wo.weight": "native-00002.safetensors",
182
+ "language_model.model.layers.31.attention.wqkv.weight": "native-00002.safetensors",
183
+ "language_model.model.layers.31.attention_norm.weight": "native-00002.safetensors",
184
+ "language_model.model.layers.31.feed_forward.w1.weight": "native-00002.safetensors",
185
+ "language_model.model.layers.31.feed_forward.w2.weight": "native-00002.safetensors",
186
+ "language_model.model.layers.31.feed_forward.w3.weight": "native-00002.safetensors",
187
+ "language_model.model.layers.31.ffn_norm.weight": "native-00002.safetensors",
188
+ "language_model.model.layers.32.attention.wo.weight": "native-00002.safetensors",
189
+ "language_model.model.layers.32.attention.wqkv.weight": "native-00002.safetensors",
190
+ "language_model.model.layers.32.attention_norm.weight": "native-00002.safetensors",
191
+ "language_model.model.layers.32.feed_forward.w1.weight": "native-00002.safetensors",
192
+ "language_model.model.layers.32.feed_forward.w2.weight": "native-00002.safetensors",
193
+ "language_model.model.layers.32.feed_forward.w3.weight": "native-00002.safetensors",
194
+ "language_model.model.layers.32.ffn_norm.weight": "native-00002.safetensors",
195
+ "language_model.model.layers.33.attention.wo.weight": "native-00002.safetensors",
196
+ "language_model.model.layers.33.attention.wqkv.weight": "native-00002.safetensors",
197
+ "language_model.model.layers.33.attention_norm.weight": "native-00002.safetensors",
198
+ "language_model.model.layers.33.feed_forward.w1.weight": "native-00002.safetensors",
199
+ "language_model.model.layers.33.feed_forward.w2.weight": "native-00002.safetensors",
200
+ "language_model.model.layers.33.feed_forward.w3.weight": "native-00002.safetensors",
201
+ "language_model.model.layers.33.ffn_norm.weight": "native-00002.safetensors",
202
+ "language_model.model.layers.34.attention.wo.weight": "native-00002.safetensors",
203
+ "language_model.model.layers.34.attention.wqkv.weight": "native-00002.safetensors",
204
+ "language_model.model.layers.34.attention_norm.weight": "native-00002.safetensors",
205
+ "language_model.model.layers.34.feed_forward.w1.weight": "native-00002.safetensors",
206
+ "language_model.model.layers.34.feed_forward.w2.weight": "native-00002.safetensors",
207
+ "language_model.model.layers.34.feed_forward.w3.weight": "native-00002.safetensors",
208
+ "language_model.model.layers.34.ffn_norm.weight": "native-00002.safetensors",
209
+ "language_model.model.layers.35.attention.wo.weight": "native-00002.safetensors",
210
+ "language_model.model.layers.35.attention.wqkv.weight": "native-00002.safetensors",
211
+ "language_model.model.layers.35.attention_norm.weight": "native-00002.safetensors",
212
+ "language_model.model.layers.35.feed_forward.w1.weight": "native-00002.safetensors",
213
+ "language_model.model.layers.35.feed_forward.w2.weight": "native-00002.safetensors",
214
+ "language_model.model.layers.35.feed_forward.w3.weight": "native-00002.safetensors",
215
+ "language_model.model.layers.35.ffn_norm.weight": "native-00002.safetensors",
216
+ "language_model.model.layers.36.attention.wo.weight": "native-00002.safetensors",
217
+ "language_model.model.layers.36.attention.wqkv.weight": "native-00002.safetensors",
218
+ "language_model.model.layers.36.attention_norm.weight": "native-00002.safetensors",
219
+ "language_model.model.layers.36.feed_forward.w1.weight": "native-00003.safetensors",
220
+ "language_model.model.layers.36.feed_forward.w2.weight": "native-00003.safetensors",
221
+ "language_model.model.layers.36.feed_forward.w3.weight": "native-00003.safetensors",
222
+ "language_model.model.layers.36.ffn_norm.weight": "native-00003.safetensors",
223
+ "language_model.model.layers.37.attention.wo.weight": "native-00003.safetensors",
224
+ "language_model.model.layers.37.attention.wqkv.weight": "native-00003.safetensors",
225
+ "language_model.model.layers.37.attention_norm.weight": "native-00003.safetensors",
226
+ "language_model.model.layers.37.feed_forward.w1.weight": "native-00003.safetensors",
227
+ "language_model.model.layers.37.feed_forward.w2.weight": "native-00003.safetensors",
228
+ "language_model.model.layers.37.feed_forward.w3.weight": "native-00003.safetensors",
229
+ "language_model.model.layers.37.ffn_norm.weight": "native-00003.safetensors",
230
+ "language_model.model.layers.38.attention.wo.weight": "native-00003.safetensors",
231
+ "language_model.model.layers.38.attention.wqkv.weight": "native-00003.safetensors",
232
+ "language_model.model.layers.38.attention_norm.weight": "native-00003.safetensors",
233
+ "language_model.model.layers.38.feed_forward.w1.weight": "native-00003.safetensors",
234
+ "language_model.model.layers.38.feed_forward.w2.weight": "native-00003.safetensors",
235
+ "language_model.model.layers.38.feed_forward.w3.weight": "native-00003.safetensors",
236
+ "language_model.model.layers.38.ffn_norm.weight": "native-00003.safetensors",
237
+ "language_model.model.layers.39.attention.wo.weight": "native-00003.safetensors",
238
+ "language_model.model.layers.39.attention.wqkv.weight": "native-00003.safetensors",
239
+ "language_model.model.layers.39.attention_norm.weight": "native-00003.safetensors",
240
+ "language_model.model.layers.39.feed_forward.w1.weight": "native-00003.safetensors",
241
+ "language_model.model.layers.39.feed_forward.w2.weight": "native-00003.safetensors",
242
+ "language_model.model.layers.39.feed_forward.w3.weight": "native-00003.safetensors",
243
+ "language_model.model.layers.39.ffn_norm.weight": "native-00003.safetensors",
244
+ "language_model.model.layers.4.attention.wo.weight": "native-00003.safetensors",
245
+ "language_model.model.layers.4.attention.wqkv.weight": "native-00003.safetensors",
246
+ "language_model.model.layers.4.attention_norm.weight": "native-00003.safetensors",
247
+ "language_model.model.layers.4.feed_forward.w1.weight": "native-00003.safetensors",
248
+ "language_model.model.layers.4.feed_forward.w2.weight": "native-00003.safetensors",
249
+ "language_model.model.layers.4.feed_forward.w3.weight": "native-00003.safetensors",
250
+ "language_model.model.layers.4.ffn_norm.weight": "native-00003.safetensors",
251
+ "language_model.model.layers.40.attention.wo.weight": "native-00003.safetensors",
252
+ "language_model.model.layers.40.attention.wqkv.weight": "native-00003.safetensors",
253
+ "language_model.model.layers.40.attention_norm.weight": "native-00003.safetensors",
254
+ "language_model.model.layers.40.feed_forward.w1.weight": "native-00003.safetensors",
255
+ "language_model.model.layers.40.feed_forward.w2.weight": "native-00003.safetensors",
256
+ "language_model.model.layers.40.feed_forward.w3.weight": "native-00003.safetensors",
257
+ "language_model.model.layers.40.ffn_norm.weight": "native-00003.safetensors",
258
+ "language_model.model.layers.41.attention.wo.weight": "native-00003.safetensors",
259
+ "language_model.model.layers.41.attention.wqkv.weight": "native-00003.safetensors",
260
+ "language_model.model.layers.41.attention_norm.weight": "native-00003.safetensors",
261
+ "language_model.model.layers.41.feed_forward.w1.weight": "native-00003.safetensors",
262
+ "language_model.model.layers.41.feed_forward.w2.weight": "native-00003.safetensors",
263
+ "language_model.model.layers.41.feed_forward.w3.weight": "native-00003.safetensors",
264
+ "language_model.model.layers.41.ffn_norm.weight": "native-00003.safetensors",
265
+ "language_model.model.layers.42.attention.wo.weight": "native-00003.safetensors",
266
+ "language_model.model.layers.42.attention.wqkv.weight": "native-00003.safetensors",
267
+ "language_model.model.layers.42.attention_norm.weight": "native-00003.safetensors",
268
+ "language_model.model.layers.42.feed_forward.w1.weight": "native-00003.safetensors",
269
+ "language_model.model.layers.42.feed_forward.w2.weight": "native-00003.safetensors",
270
+ "language_model.model.layers.42.feed_forward.w3.weight": "native-00003.safetensors",
271
+ "language_model.model.layers.42.ffn_norm.weight": "native-00003.safetensors",
272
+ "language_model.model.layers.43.attention.wo.weight": "native-00003.safetensors",
273
+ "language_model.model.layers.43.attention.wqkv.weight": "native-00003.safetensors",
274
+ "language_model.model.layers.43.attention_norm.weight": "native-00003.safetensors",
275
+ "language_model.model.layers.43.feed_forward.w1.weight": "native-00003.safetensors",
276
+ "language_model.model.layers.43.feed_forward.w2.weight": "native-00003.safetensors",
277
+ "language_model.model.layers.43.feed_forward.w3.weight": "native-00003.safetensors",
278
+ "language_model.model.layers.43.ffn_norm.weight": "native-00003.safetensors",
279
+ "language_model.model.layers.44.attention.wo.weight": "native-00003.safetensors",
280
+ "language_model.model.layers.44.attention.wqkv.weight": "native-00003.safetensors",
281
+ "language_model.model.layers.44.attention_norm.weight": "native-00003.safetensors",
282
+ "language_model.model.layers.44.feed_forward.w1.weight": "native-00003.safetensors",
283
+ "language_model.model.layers.44.feed_forward.w2.weight": "native-00003.safetensors",
284
+ "language_model.model.layers.44.feed_forward.w3.weight": "native-00003.safetensors",
285
+ "language_model.model.layers.44.ffn_norm.weight": "native-00003.safetensors",
286
+ "language_model.model.layers.45.attention.wo.weight": "native-00003.safetensors",
287
+ "language_model.model.layers.45.attention.wqkv.weight": "native-00003.safetensors",
288
+ "language_model.model.layers.45.attention_norm.weight": "native-00003.safetensors",
289
+ "language_model.model.layers.45.feed_forward.w1.weight": "native-00003.safetensors",
290
+ "language_model.model.layers.45.feed_forward.w2.weight": "native-00003.safetensors",
291
+ "language_model.model.layers.45.feed_forward.w3.weight": "native-00003.safetensors",
292
+ "language_model.model.layers.45.ffn_norm.weight": "native-00003.safetensors",
293
+ "language_model.model.layers.46.attention.wo.weight": "native-00003.safetensors",
294
+ "language_model.model.layers.46.attention.wqkv.weight": "native-00003.safetensors",
295
+ "language_model.model.layers.46.attention_norm.weight": "native-00003.safetensors",
296
+ "language_model.model.layers.46.feed_forward.w1.weight": "native-00003.safetensors",
297
+ "language_model.model.layers.46.feed_forward.w2.weight": "native-00003.safetensors",
298
+ "language_model.model.layers.46.feed_forward.w3.weight": "native-00003.safetensors",
299
+ "language_model.model.layers.46.ffn_norm.weight": "native-00003.safetensors",
300
+ "language_model.model.layers.47.attention.wo.weight": "native-00003.safetensors",
301
+ "language_model.model.layers.47.attention.wqkv.weight": "native-00003.safetensors",
302
+ "language_model.model.layers.47.attention_norm.weight": "native-00003.safetensors",
303
+ "language_model.model.layers.47.feed_forward.w1.weight": "native-00003.safetensors",
304
+ "language_model.model.layers.47.feed_forward.w2.weight": "native-00003.safetensors",
305
+ "language_model.model.layers.47.feed_forward.w3.weight": "native-00003.safetensors",
306
+ "language_model.model.layers.47.ffn_norm.weight": "native-00003.safetensors",
307
+ "language_model.model.layers.5.attention.wo.weight": "native-00003.safetensors",
308
+ "language_model.model.layers.5.attention.wqkv.weight": "native-00003.safetensors",
309
+ "language_model.model.layers.5.attention_norm.weight": "native-00003.safetensors",
310
+ "language_model.model.layers.5.feed_forward.w1.weight": "native-00003.safetensors",
311
+ "language_model.model.layers.5.feed_forward.w2.weight": "native-00003.safetensors",
312
+ "language_model.model.layers.5.feed_forward.w3.weight": "native-00003.safetensors",
313
+ "language_model.model.layers.5.ffn_norm.weight": "native-00003.safetensors",
314
+ "language_model.model.layers.6.attention.wo.weight": "native-00003.safetensors",
315
+ "language_model.model.layers.6.attention.wqkv.weight": "native-00003.safetensors",
316
+ "language_model.model.layers.6.attention_norm.weight": "native-00003.safetensors",
317
+ "language_model.model.layers.6.feed_forward.w1.weight": "native-00003.safetensors",
318
+ "language_model.model.layers.6.feed_forward.w2.weight": "native-00003.safetensors",
319
+ "language_model.model.layers.6.feed_forward.w3.weight": "native-00003.safetensors",
320
+ "language_model.model.layers.6.ffn_norm.weight": "native-00003.safetensors",
321
+ "language_model.model.layers.7.attention.wo.weight": "native-00003.safetensors",
322
+ "language_model.model.layers.7.attention.wqkv.weight": "native-00003.safetensors",
323
+ "language_model.model.layers.7.attention_norm.weight": "native-00003.safetensors",
324
+ "language_model.model.layers.7.feed_forward.w1.weight": "native-00004.safetensors",
325
+ "language_model.model.layers.7.feed_forward.w2.weight": "native-00004.safetensors",
326
+ "language_model.model.layers.7.feed_forward.w3.weight": "native-00004.safetensors",
327
+ "language_model.model.layers.7.ffn_norm.weight": "native-00004.safetensors",
328
+ "language_model.model.layers.8.attention.wo.weight": "native-00004.safetensors",
329
+ "language_model.model.layers.8.attention.wqkv.weight": "native-00004.safetensors",
330
+ "language_model.model.layers.8.attention_norm.weight": "native-00004.safetensors",
331
+ "language_model.model.layers.8.feed_forward.w1.weight": "native-00004.safetensors",
332
+ "language_model.model.layers.8.feed_forward.w2.weight": "native-00004.safetensors",
333
+ "language_model.model.layers.8.feed_forward.w3.weight": "native-00004.safetensors",
334
+ "language_model.model.layers.8.ffn_norm.weight": "native-00004.safetensors",
335
+ "language_model.model.layers.9.attention.wo.weight": "native-00004.safetensors",
336
+ "language_model.model.layers.9.attention.wqkv.weight": "native-00004.safetensors",
337
+ "language_model.model.layers.9.attention_norm.weight": "native-00004.safetensors",
338
+ "language_model.model.layers.9.feed_forward.w1.weight": "native-00004.safetensors",
339
+ "language_model.model.layers.9.feed_forward.w2.weight": "native-00004.safetensors",
340
+ "language_model.model.layers.9.feed_forward.w3.weight": "native-00004.safetensors",
341
+ "language_model.model.layers.9.ffn_norm.weight": "native-00004.safetensors",
342
+ "language_model.model.norm.weight": "native-00004.safetensors",
343
+ "language_model.model.tok_embeddings.weight": "native-00004.safetensors",
344
+ "language_model.output.weight": "native-00004.safetensors",
345
+ "mlp1.0.bias": "native-00004.safetensors",
346
+ "mlp1.0.weight": "native-00004.safetensors",
347
+ "mlp1.1.bias": "native-00004.safetensors",
348
+ "mlp1.1.weight": "native-00004.safetensors",
349
+ "mlp1.3.bias": "native-00004.safetensors",
350
+ "mlp1.3.weight": "native-00004.safetensors",
351
+ "vision_model.embeddings.class_embedding": "native-00004.safetensors",
352
+ "vision_model.embeddings.patch_embedding.bias": "native-00004.safetensors",
353
+ "vision_model.embeddings.patch_embedding.weight": "native-00004.safetensors",
354
+ "vision_model.embeddings.position_embedding": "native-00004.safetensors",
355
+ "vision_model.encoder.layers.0.attn.proj.bias": "native-00004.safetensors",
356
+ "vision_model.encoder.layers.0.attn.proj.weight": "native-00004.safetensors",
357
+ "vision_model.encoder.layers.0.attn.qkv.bias": "native-00004.safetensors",
358
+ "vision_model.encoder.layers.0.attn.qkv.weight": "native-00004.safetensors",
359
+ "vision_model.encoder.layers.0.ls1": "native-00004.safetensors",
360
+ "vision_model.encoder.layers.0.ls2": "native-00004.safetensors",
361
+ "vision_model.encoder.layers.0.mlp.fc1.bias": "native-00004.safetensors",
362
+ "vision_model.encoder.layers.0.mlp.fc1.weight": "native-00004.safetensors",
363
+ "vision_model.encoder.layers.0.mlp.fc2.bias": "native-00004.safetensors",
364
+ "vision_model.encoder.layers.0.mlp.fc2.weight": "native-00004.safetensors",
365
+ "vision_model.encoder.layers.0.norm1.bias": "native-00004.safetensors",
366
+ "vision_model.encoder.layers.0.norm1.weight": "native-00004.safetensors",
367
+ "vision_model.encoder.layers.0.norm2.bias": "native-00004.safetensors",
368
+ "vision_model.encoder.layers.0.norm2.weight": "native-00004.safetensors",
369
+ "vision_model.encoder.layers.1.attn.proj.bias": "native-00004.safetensors",
370
+ "vision_model.encoder.layers.1.attn.proj.weight": "native-00004.safetensors",
371
+ "vision_model.encoder.layers.1.attn.qkv.bias": "native-00004.safetensors",
372
+ "vision_model.encoder.layers.1.attn.qkv.weight": "native-00004.safetensors",
373
+ "vision_model.encoder.layers.1.ls1": "native-00004.safetensors",
374
+ "vision_model.encoder.layers.1.ls2": "native-00004.safetensors",
375
+ "vision_model.encoder.layers.1.mlp.fc1.bias": "native-00004.safetensors",
376
+ "vision_model.encoder.layers.1.mlp.fc1.weight": "native-00004.safetensors",
377
+ "vision_model.encoder.layers.1.mlp.fc2.bias": "native-00004.safetensors",
378
+ "vision_model.encoder.layers.1.mlp.fc2.weight": "native-00004.safetensors",
379
+ "vision_model.encoder.layers.1.norm1.bias": "native-00004.safetensors",
380
+ "vision_model.encoder.layers.1.norm1.weight": "native-00004.safetensors",
381
+ "vision_model.encoder.layers.1.norm2.bias": "native-00004.safetensors",
382
+ "vision_model.encoder.layers.1.norm2.weight": "native-00004.safetensors",
383
+ "vision_model.encoder.layers.10.attn.proj.bias": "native-00004.safetensors",
384
+ "vision_model.encoder.layers.10.attn.proj.weight": "native-00004.safetensors",
385
+ "vision_model.encoder.layers.10.attn.qkv.bias": "native-00004.safetensors",
386
+ "vision_model.encoder.layers.10.attn.qkv.weight": "native-00004.safetensors",
387
+ "vision_model.encoder.layers.10.ls1": "native-00004.safetensors",
388
+ "vision_model.encoder.layers.10.ls2": "native-00004.safetensors",
389
+ "vision_model.encoder.layers.10.mlp.fc1.bias": "native-00004.safetensors",
390
+ "vision_model.encoder.layers.10.mlp.fc1.weight": "native-00004.safetensors",
391
+ "vision_model.encoder.layers.10.mlp.fc2.bias": "native-00004.safetensors",
392
+ "vision_model.encoder.layers.10.mlp.fc2.weight": "native-00004.safetensors",
393
+ "vision_model.encoder.layers.10.norm1.bias": "native-00004.safetensors",
394
+ "vision_model.encoder.layers.10.norm1.weight": "native-00004.safetensors",
395
+ "vision_model.encoder.layers.10.norm2.bias": "native-00004.safetensors",
396
+ "vision_model.encoder.layers.10.norm2.weight": "native-00004.safetensors",
397
+ "vision_model.encoder.layers.11.attn.proj.bias": "native-00004.safetensors",
398
+ "vision_model.encoder.layers.11.attn.proj.weight": "native-00004.safetensors",
399
+ "vision_model.encoder.layers.11.attn.qkv.bias": "native-00004.safetensors",
400
+ "vision_model.encoder.layers.11.attn.qkv.weight": "native-00004.safetensors",
401
+ "vision_model.encoder.layers.11.ls1": "native-00004.safetensors",
402
+ "vision_model.encoder.layers.11.ls2": "native-00004.safetensors",
403
+ "vision_model.encoder.layers.11.mlp.fc1.bias": "native-00004.safetensors",
404
+ "vision_model.encoder.layers.11.mlp.fc1.weight": "native-00004.safetensors",
405
+ "vision_model.encoder.layers.11.mlp.fc2.bias": "native-00004.safetensors",
406
+ "vision_model.encoder.layers.11.mlp.fc2.weight": "native-00004.safetensors",
407
+ "vision_model.encoder.layers.11.norm1.bias": "native-00004.safetensors",
408
+ "vision_model.encoder.layers.11.norm1.weight": "native-00004.safetensors",
409
+ "vision_model.encoder.layers.11.norm2.bias": "native-00004.safetensors",
410
+ "vision_model.encoder.layers.11.norm2.weight": "native-00004.safetensors",
411
+ "vision_model.encoder.layers.12.attn.proj.bias": "native-00004.safetensors",
412
+ "vision_model.encoder.layers.12.attn.proj.weight": "native-00004.safetensors",
413
+ "vision_model.encoder.layers.12.attn.qkv.bias": "native-00004.safetensors",
414
+ "vision_model.encoder.layers.12.attn.qkv.weight": "native-00004.safetensors",
415
+ "vision_model.encoder.layers.12.ls1": "native-00004.safetensors",
416
+ "vision_model.encoder.layers.12.ls2": "native-00004.safetensors",
417
+ "vision_model.encoder.layers.12.mlp.fc1.bias": "native-00004.safetensors",
418
+ "vision_model.encoder.layers.12.mlp.fc1.weight": "native-00004.safetensors",
419
+ "vision_model.encoder.layers.12.mlp.fc2.bias": "native-00004.safetensors",
420
+ "vision_model.encoder.layers.12.mlp.fc2.weight": "native-00004.safetensors",
421
+ "vision_model.encoder.layers.12.norm1.bias": "native-00004.safetensors",
422
+ "vision_model.encoder.layers.12.norm1.weight": "native-00004.safetensors",
423
+ "vision_model.encoder.layers.12.norm2.bias": "native-00004.safetensors",
424
+ "vision_model.encoder.layers.12.norm2.weight": "native-00004.safetensors",
425
+ "vision_model.encoder.layers.13.attn.proj.bias": "native-00004.safetensors",
426
+ "vision_model.encoder.layers.13.attn.proj.weight": "native-00004.safetensors",
427
+ "vision_model.encoder.layers.13.attn.qkv.bias": "native-00004.safetensors",
428
+ "vision_model.encoder.layers.13.attn.qkv.weight": "native-00004.safetensors",
429
+ "vision_model.encoder.layers.13.ls1": "native-00004.safetensors",
430
+ "vision_model.encoder.layers.13.ls2": "native-00004.safetensors",
431
+ "vision_model.encoder.layers.13.mlp.fc1.bias": "native-00004.safetensors",
432
+ "vision_model.encoder.layers.13.mlp.fc1.weight": "native-00004.safetensors",
433
+ "vision_model.encoder.layers.13.mlp.fc2.bias": "native-00004.safetensors",
434
+ "vision_model.encoder.layers.13.mlp.fc2.weight": "native-00004.safetensors",
435
+ "vision_model.encoder.layers.13.norm1.bias": "native-00004.safetensors",
436
+ "vision_model.encoder.layers.13.norm1.weight": "native-00004.safetensors",
437
+ "vision_model.encoder.layers.13.norm2.bias": "native-00004.safetensors",
438
+ "vision_model.encoder.layers.13.norm2.weight": "native-00004.safetensors",
439
+ "vision_model.encoder.layers.14.attn.proj.bias": "native-00004.safetensors",
440
+ "vision_model.encoder.layers.14.attn.proj.weight": "native-00004.safetensors",
441
+ "vision_model.encoder.layers.14.attn.qkv.bias": "native-00004.safetensors",
442
+ "vision_model.encoder.layers.14.attn.qkv.weight": "native-00004.safetensors",
443
+ "vision_model.encoder.layers.14.ls1": "native-00004.safetensors",
444
+ "vision_model.encoder.layers.14.ls2": "native-00004.safetensors",
445
+ "vision_model.encoder.layers.14.mlp.fc1.bias": "native-00004.safetensors",
446
+ "vision_model.encoder.layers.14.mlp.fc1.weight": "native-00004.safetensors",
447
+ "vision_model.encoder.layers.14.mlp.fc2.bias": "native-00004.safetensors",
448
+ "vision_model.encoder.layers.14.mlp.fc2.weight": "native-00004.safetensors",
449
+ "vision_model.encoder.layers.14.norm1.bias": "native-00004.safetensors",
450
+ "vision_model.encoder.layers.14.norm1.weight": "native-00004.safetensors",
451
+ "vision_model.encoder.layers.14.norm2.bias": "native-00004.safetensors",
452
+ "vision_model.encoder.layers.14.norm2.weight": "native-00004.safetensors",
453
+ "vision_model.encoder.layers.15.attn.proj.bias": "native-00004.safetensors",
454
+ "vision_model.encoder.layers.15.attn.proj.weight": "native-00004.safetensors",
455
+ "vision_model.encoder.layers.15.attn.qkv.bias": "native-00004.safetensors",
456
+ "vision_model.encoder.layers.15.attn.qkv.weight": "native-00004.safetensors",
457
+ "vision_model.encoder.layers.15.ls1": "native-00004.safetensors",
458
+ "vision_model.encoder.layers.15.ls2": "native-00004.safetensors",
459
+ "vision_model.encoder.layers.15.mlp.fc1.bias": "native-00004.safetensors",
460
+ "vision_model.encoder.layers.15.mlp.fc1.weight": "native-00004.safetensors",
461
+ "vision_model.encoder.layers.15.mlp.fc2.bias": "native-00004.safetensors",
462
+ "vision_model.encoder.layers.15.mlp.fc2.weight": "native-00004.safetensors",
463
+ "vision_model.encoder.layers.15.norm1.bias": "native-00004.safetensors",
464
+ "vision_model.encoder.layers.15.norm1.weight": "native-00004.safetensors",
465
+ "vision_model.encoder.layers.15.norm2.bias": "native-00004.safetensors",
466
+ "vision_model.encoder.layers.15.norm2.weight": "native-00004.safetensors",
467
+ "vision_model.encoder.layers.16.attn.proj.bias": "native-00004.safetensors",
468
+ "vision_model.encoder.layers.16.attn.proj.weight": "native-00004.safetensors",
469
+ "vision_model.encoder.layers.16.attn.qkv.bias": "native-00004.safetensors",
470
+ "vision_model.encoder.layers.16.attn.qkv.weight": "native-00004.safetensors",
471
+ "vision_model.encoder.layers.16.ls1": "native-00004.safetensors",
472
+ "vision_model.encoder.layers.16.ls2": "native-00004.safetensors",
473
+ "vision_model.encoder.layers.16.mlp.fc1.bias": "native-00004.safetensors",
474
+ "vision_model.encoder.layers.16.mlp.fc1.weight": "native-00004.safetensors",
475
+ "vision_model.encoder.layers.16.mlp.fc2.bias": "native-00004.safetensors",
476
+ "vision_model.encoder.layers.16.mlp.fc2.weight": "native-00004.safetensors",
477
+ "vision_model.encoder.layers.16.norm1.bias": "native-00004.safetensors",
478
+ "vision_model.encoder.layers.16.norm1.weight": "native-00004.safetensors",
479
+ "vision_model.encoder.layers.16.norm2.bias": "native-00004.safetensors",
480
+ "vision_model.encoder.layers.16.norm2.weight": "native-00004.safetensors",
481
+ "vision_model.encoder.layers.17.attn.proj.bias": "native-00004.safetensors",
482
+ "vision_model.encoder.layers.17.attn.proj.weight": "native-00004.safetensors",
483
+ "vision_model.encoder.layers.17.attn.qkv.bias": "native-00004.safetensors",
484
+ "vision_model.encoder.layers.17.attn.qkv.weight": "native-00004.safetensors",
485
+ "vision_model.encoder.layers.17.ls1": "native-00004.safetensors",
486
+ "vision_model.encoder.layers.17.ls2": "native-00004.safetensors",
487
+ "vision_model.encoder.layers.17.mlp.fc1.bias": "native-00004.safetensors",
488
+ "vision_model.encoder.layers.17.mlp.fc1.weight": "native-00004.safetensors",
489
+ "vision_model.encoder.layers.17.mlp.fc2.bias": "native-00004.safetensors",
490
+ "vision_model.encoder.layers.17.mlp.fc2.weight": "native-00004.safetensors",
491
+ "vision_model.encoder.layers.17.norm1.bias": "native-00004.safetensors",
492
+ "vision_model.encoder.layers.17.norm1.weight": "native-00004.safetensors",
493
+ "vision_model.encoder.layers.17.norm2.bias": "native-00004.safetensors",
494
+ "vision_model.encoder.layers.17.norm2.weight": "native-00004.safetensors",
495
+ "vision_model.encoder.layers.18.attn.proj.bias": "native-00004.safetensors",
496
+ "vision_model.encoder.layers.18.attn.proj.weight": "native-00004.safetensors",
497
+ "vision_model.encoder.layers.18.attn.qkv.bias": "native-00004.safetensors",
498
+ "vision_model.encoder.layers.18.attn.qkv.weight": "native-00004.safetensors",
499
+ "vision_model.encoder.layers.18.ls1": "native-00004.safetensors",
500
+ "vision_model.encoder.layers.18.ls2": "native-00004.safetensors",
501
+ "vision_model.encoder.layers.18.mlp.fc1.bias": "native-00004.safetensors",
502
+ "vision_model.encoder.layers.18.mlp.fc1.weight": "native-00004.safetensors",
503
+ "vision_model.encoder.layers.18.mlp.fc2.bias": "native-00004.safetensors",
504
+ "vision_model.encoder.layers.18.mlp.fc2.weight": "native-00004.safetensors",
505
+ "vision_model.encoder.layers.18.norm1.bias": "native-00004.safetensors",
506
+ "vision_model.encoder.layers.18.norm1.weight": "native-00004.safetensors",
507
+ "vision_model.encoder.layers.18.norm2.bias": "native-00004.safetensors",
508
+ "vision_model.encoder.layers.18.norm2.weight": "native-00004.safetensors",
509
+ "vision_model.encoder.layers.19.attn.proj.bias": "native-00004.safetensors",
510
+ "vision_model.encoder.layers.19.attn.proj.weight": "native-00004.safetensors",
511
+ "vision_model.encoder.layers.19.attn.qkv.bias": "native-00004.safetensors",
512
+ "vision_model.encoder.layers.19.attn.qkv.weight": "native-00004.safetensors",
513
+ "vision_model.encoder.layers.19.ls1": "native-00004.safetensors",
514
+ "vision_model.encoder.layers.19.ls2": "native-00004.safetensors",
515
+ "vision_model.encoder.layers.19.mlp.fc1.bias": "native-00004.safetensors",
516
+ "vision_model.encoder.layers.19.mlp.fc1.weight": "native-00004.safetensors",
517
+ "vision_model.encoder.layers.19.mlp.fc2.bias": "native-00004.safetensors",
518
+ "vision_model.encoder.layers.19.mlp.fc2.weight": "native-00004.safetensors",
519
+ "vision_model.encoder.layers.19.norm1.bias": "native-00004.safetensors",
520
+ "vision_model.encoder.layers.19.norm1.weight": "native-00004.safetensors",
521
+ "vision_model.encoder.layers.19.norm2.bias": "native-00004.safetensors",
522
+ "vision_model.encoder.layers.19.norm2.weight": "native-00004.safetensors",
523
+ "vision_model.encoder.layers.2.attn.proj.bias": "native-00004.safetensors",
524
+ "vision_model.encoder.layers.2.attn.proj.weight": "native-00004.safetensors",
525
+ "vision_model.encoder.layers.2.attn.qkv.bias": "native-00004.safetensors",
526
+ "vision_model.encoder.layers.2.attn.qkv.weight": "native-00004.safetensors",
527
+ "vision_model.encoder.layers.2.ls1": "native-00004.safetensors",
528
+ "vision_model.encoder.layers.2.ls2": "native-00004.safetensors",
529
+ "vision_model.encoder.layers.2.mlp.fc1.bias": "native-00004.safetensors",
530
+ "vision_model.encoder.layers.2.mlp.fc1.weight": "native-00004.safetensors",
531
+ "vision_model.encoder.layers.2.mlp.fc2.bias": "native-00004.safetensors",
532
+ "vision_model.encoder.layers.2.mlp.fc2.weight": "native-00004.safetensors",
533
+ "vision_model.encoder.layers.2.norm1.bias": "native-00004.safetensors",
534
+ "vision_model.encoder.layers.2.norm1.weight": "native-00004.safetensors",
535
+ "vision_model.encoder.layers.2.norm2.bias": "native-00004.safetensors",
536
+ "vision_model.encoder.layers.2.norm2.weight": "native-00004.safetensors",
537
+ "vision_model.encoder.layers.20.attn.proj.bias": "native-00004.safetensors",
538
+ "vision_model.encoder.layers.20.attn.proj.weight": "native-00004.safetensors",
539
+ "vision_model.encoder.layers.20.attn.qkv.bias": "native-00004.safetensors",
540
+ "vision_model.encoder.layers.20.attn.qkv.weight": "native-00004.safetensors",
541
+ "vision_model.encoder.layers.20.ls1": "native-00004.safetensors",
542
+ "vision_model.encoder.layers.20.ls2": "native-00004.safetensors",
543
+ "vision_model.encoder.layers.20.mlp.fc1.bias": "native-00004.safetensors",
544
+ "vision_model.encoder.layers.20.mlp.fc1.weight": "native-00004.safetensors",
545
+ "vision_model.encoder.layers.20.mlp.fc2.bias": "native-00004.safetensors",
546
+ "vision_model.encoder.layers.20.mlp.fc2.weight": "native-00004.safetensors",
547
+ "vision_model.encoder.layers.20.norm1.bias": "native-00004.safetensors",
548
+ "vision_model.encoder.layers.20.norm1.weight": "native-00004.safetensors",
549
+ "vision_model.encoder.layers.20.norm2.bias": "native-00004.safetensors",
550
+ "vision_model.encoder.layers.20.norm2.weight": "native-00004.safetensors",
551
+ "vision_model.encoder.layers.21.attn.proj.bias": "native-00004.safetensors",
552
+ "vision_model.encoder.layers.21.attn.proj.weight": "native-00004.safetensors",
553
+ "vision_model.encoder.layers.21.attn.qkv.bias": "native-00004.safetensors",
554
+ "vision_model.encoder.layers.21.attn.qkv.weight": "native-00004.safetensors",
555
+ "vision_model.encoder.layers.21.ls1": "native-00004.safetensors",
556
+ "vision_model.encoder.layers.21.ls2": "native-00004.safetensors",
557
+ "vision_model.encoder.layers.21.mlp.fc1.bias": "native-00004.safetensors",
558
+ "vision_model.encoder.layers.21.mlp.fc1.weight": "native-00004.safetensors",
559
+ "vision_model.encoder.layers.21.mlp.fc2.bias": "native-00004.safetensors",
560
+ "vision_model.encoder.layers.21.mlp.fc2.weight": "native-00004.safetensors",
561
+ "vision_model.encoder.layers.21.norm1.bias": "native-00004.safetensors",
562
+ "vision_model.encoder.layers.21.norm1.weight": "native-00004.safetensors",
563
+ "vision_model.encoder.layers.21.norm2.bias": "native-00004.safetensors",
564
+ "vision_model.encoder.layers.21.norm2.weight": "native-00004.safetensors",
565
+ "vision_model.encoder.layers.22.attn.proj.bias": "native-00004.safetensors",
566
+ "vision_model.encoder.layers.22.attn.proj.weight": "native-00004.safetensors",
567
+ "vision_model.encoder.layers.22.attn.qkv.bias": "native-00004.safetensors",
568
+ "vision_model.encoder.layers.22.attn.qkv.weight": "native-00004.safetensors",
569
+ "vision_model.encoder.layers.22.ls1": "native-00004.safetensors",
570
+ "vision_model.encoder.layers.22.ls2": "native-00004.safetensors",
571
+ "vision_model.encoder.layers.22.mlp.fc1.bias": "native-00004.safetensors",
572
+ "vision_model.encoder.layers.22.mlp.fc1.weight": "native-00004.safetensors",
573
+ "vision_model.encoder.layers.22.mlp.fc2.bias": "native-00004.safetensors",
574
+ "vision_model.encoder.layers.22.mlp.fc2.weight": "native-00004.safetensors",
575
+ "vision_model.encoder.layers.22.norm1.bias": "native-00004.safetensors",
576
+ "vision_model.encoder.layers.22.norm1.weight": "native-00004.safetensors",
577
+ "vision_model.encoder.layers.22.norm2.bias": "native-00004.safetensors",
578
+ "vision_model.encoder.layers.22.norm2.weight": "native-00004.safetensors",
579
+ "vision_model.encoder.layers.23.attn.proj.bias": "native-00004.safetensors",
580
+ "vision_model.encoder.layers.23.attn.proj.weight": "native-00004.safetensors",
581
+ "vision_model.encoder.layers.23.attn.qkv.bias": "native-00004.safetensors",
582
+ "vision_model.encoder.layers.23.attn.qkv.weight": "native-00004.safetensors",
583
+ "vision_model.encoder.layers.23.ls1": "native-00004.safetensors",
584
+ "vision_model.encoder.layers.23.ls2": "native-00004.safetensors",
585
+ "vision_model.encoder.layers.23.mlp.fc1.bias": "native-00004.safetensors",
586
+ "vision_model.encoder.layers.23.mlp.fc1.weight": "native-00004.safetensors",
587
+ "vision_model.encoder.layers.23.mlp.fc2.bias": "native-00004.safetensors",
588
+ "vision_model.encoder.layers.23.mlp.fc2.weight": "native-00004.safetensors",
589
+ "vision_model.encoder.layers.23.norm1.bias": "native-00004.safetensors",
590
+ "vision_model.encoder.layers.23.norm1.weight": "native-00004.safetensors",
591
+ "vision_model.encoder.layers.23.norm2.bias": "native-00004.safetensors",
592
+ "vision_model.encoder.layers.23.norm2.weight": "native-00004.safetensors",
593
+ "vision_model.encoder.layers.3.attn.proj.bias": "native-00004.safetensors",
594
+ "vision_model.encoder.layers.3.attn.proj.weight": "native-00004.safetensors",
595
+ "vision_model.encoder.layers.3.attn.qkv.bias": "native-00004.safetensors",
596
+ "vision_model.encoder.layers.3.attn.qkv.weight": "native-00004.safetensors",
597
+ "vision_model.encoder.layers.3.ls1": "native-00004.safetensors",
598
+ "vision_model.encoder.layers.3.ls2": "native-00004.safetensors",
599
+ "vision_model.encoder.layers.3.mlp.fc1.bias": "native-00004.safetensors",
600
+ "vision_model.encoder.layers.3.mlp.fc1.weight": "native-00004.safetensors",
601
+ "vision_model.encoder.layers.3.mlp.fc2.bias": "native-00004.safetensors",
602
+ "vision_model.encoder.layers.3.mlp.fc2.weight": "native-00004.safetensors",
603
+ "vision_model.encoder.layers.3.norm1.bias": "native-00004.safetensors",
604
+ "vision_model.encoder.layers.3.norm1.weight": "native-00004.safetensors",
605
+ "vision_model.encoder.layers.3.norm2.bias": "native-00004.safetensors",
606
+ "vision_model.encoder.layers.3.norm2.weight": "native-00004.safetensors",
607
+ "vision_model.encoder.layers.4.attn.proj.bias": "native-00004.safetensors",
608
+ "vision_model.encoder.layers.4.attn.proj.weight": "native-00004.safetensors",
609
+ "vision_model.encoder.layers.4.attn.qkv.bias": "native-00004.safetensors",
610
+ "vision_model.encoder.layers.4.attn.qkv.weight": "native-00004.safetensors",
611
+ "vision_model.encoder.layers.4.ls1": "native-00004.safetensors",
612
+ "vision_model.encoder.layers.4.ls2": "native-00004.safetensors",
613
+ "vision_model.encoder.layers.4.mlp.fc1.bias": "native-00004.safetensors",
614
+ "vision_model.encoder.layers.4.mlp.fc1.weight": "native-00004.safetensors",
615
+ "vision_model.encoder.layers.4.mlp.fc2.bias": "native-00004.safetensors",
616
+ "vision_model.encoder.layers.4.mlp.fc2.weight": "native-00004.safetensors",
617
+ "vision_model.encoder.layers.4.norm1.bias": "native-00004.safetensors",
618
+ "vision_model.encoder.layers.4.norm1.weight": "native-00004.safetensors",
619
+ "vision_model.encoder.layers.4.norm2.bias": "native-00004.safetensors",
620
+ "vision_model.encoder.layers.4.norm2.weight": "native-00004.safetensors",
621
+ "vision_model.encoder.layers.5.attn.proj.bias": "native-00004.safetensors",
622
+ "vision_model.encoder.layers.5.attn.proj.weight": "native-00004.safetensors",
623
+ "vision_model.encoder.layers.5.attn.qkv.bias": "native-00004.safetensors",
624
+ "vision_model.encoder.layers.5.attn.qkv.weight": "native-00004.safetensors",
625
+ "vision_model.encoder.layers.5.ls1": "native-00004.safetensors",
626
+ "vision_model.encoder.layers.5.ls2": "native-00004.safetensors",
627
+ "vision_model.encoder.layers.5.mlp.fc1.bias": "native-00004.safetensors",
628
+ "vision_model.encoder.layers.5.mlp.fc1.weight": "native-00004.safetensors",
629
+ "vision_model.encoder.layers.5.mlp.fc2.bias": "native-00004.safetensors",
630
+ "vision_model.encoder.layers.5.mlp.fc2.weight": "native-00004.safetensors",
631
+ "vision_model.encoder.layers.5.norm1.bias": "native-00004.safetensors",
632
+ "vision_model.encoder.layers.5.norm1.weight": "native-00004.safetensors",
633
+ "vision_model.encoder.layers.5.norm2.bias": "native-00004.safetensors",
634
+ "vision_model.encoder.layers.5.norm2.weight": "native-00004.safetensors",
635
+ "vision_model.encoder.layers.6.attn.proj.bias": "native-00004.safetensors",
636
+ "vision_model.encoder.layers.6.attn.proj.weight": "native-00004.safetensors",
637
+ "vision_model.encoder.layers.6.attn.qkv.bias": "native-00004.safetensors",
638
+ "vision_model.encoder.layers.6.attn.qkv.weight": "native-00004.safetensors",
639
+ "vision_model.encoder.layers.6.ls1": "native-00004.safetensors",
640
+ "vision_model.encoder.layers.6.ls2": "native-00004.safetensors",
641
+ "vision_model.encoder.layers.6.mlp.fc1.bias": "native-00004.safetensors",
642
+ "vision_model.encoder.layers.6.mlp.fc1.weight": "native-00004.safetensors",
643
+ "vision_model.encoder.layers.6.mlp.fc2.bias": "native-00004.safetensors",
644
+ "vision_model.encoder.layers.6.mlp.fc2.weight": "native-00004.safetensors",
645
+ "vision_model.encoder.layers.6.norm1.bias": "native-00004.safetensors",
646
+ "vision_model.encoder.layers.6.norm1.weight": "native-00004.safetensors",
647
+ "vision_model.encoder.layers.6.norm2.bias": "native-00004.safetensors",
648
+ "vision_model.encoder.layers.6.norm2.weight": "native-00004.safetensors",
649
+ "vision_model.encoder.layers.7.attn.proj.bias": "native-00004.safetensors",
650
+ "vision_model.encoder.layers.7.attn.proj.weight": "native-00004.safetensors",
651
+ "vision_model.encoder.layers.7.attn.qkv.bias": "native-00004.safetensors",
652
+ "vision_model.encoder.layers.7.attn.qkv.weight": "native-00004.safetensors",
653
+ "vision_model.encoder.layers.7.ls1": "native-00004.safetensors",
654
+ "vision_model.encoder.layers.7.ls2": "native-00004.safetensors",
655
+ "vision_model.encoder.layers.7.mlp.fc1.bias": "native-00004.safetensors",
656
+ "vision_model.encoder.layers.7.mlp.fc1.weight": "native-00004.safetensors",
657
+ "vision_model.encoder.layers.7.mlp.fc2.bias": "native-00004.safetensors",
658
+ "vision_model.encoder.layers.7.mlp.fc2.weight": "native-00004.safetensors",
659
+ "vision_model.encoder.layers.7.norm1.bias": "native-00004.safetensors",
660
+ "vision_model.encoder.layers.7.norm1.weight": "native-00004.safetensors",
661
+ "vision_model.encoder.layers.7.norm2.bias": "native-00004.safetensors",
662
+ "vision_model.encoder.layers.7.norm2.weight": "native-00004.safetensors",
663
+ "vision_model.encoder.layers.8.attn.proj.bias": "native-00004.safetensors",
664
+ "vision_model.encoder.layers.8.attn.proj.weight": "native-00004.safetensors",
665
+ "vision_model.encoder.layers.8.attn.qkv.bias": "native-00004.safetensors",
666
+ "vision_model.encoder.layers.8.attn.qkv.weight": "native-00004.safetensors",
667
+ "vision_model.encoder.layers.8.ls1": "native-00004.safetensors",
668
+ "vision_model.encoder.layers.8.ls2": "native-00004.safetensors",
669
+ "vision_model.encoder.layers.8.mlp.fc1.bias": "native-00004.safetensors",
670
+ "vision_model.encoder.layers.8.mlp.fc1.weight": "native-00004.safetensors",
671
+ "vision_model.encoder.layers.8.mlp.fc2.bias": "native-00004.safetensors",
672
+ "vision_model.encoder.layers.8.mlp.fc2.weight": "native-00004.safetensors",
673
+ "vision_model.encoder.layers.8.norm1.bias": "native-00004.safetensors",
674
+ "vision_model.encoder.layers.8.norm1.weight": "native-00004.safetensors",
675
+ "vision_model.encoder.layers.8.norm2.bias": "native-00004.safetensors",
676
+ "vision_model.encoder.layers.8.norm2.weight": "native-00004.safetensors",
677
+ "vision_model.encoder.layers.9.attn.proj.bias": "native-00004.safetensors",
678
+ "vision_model.encoder.layers.9.attn.proj.weight": "native-00004.safetensors",
679
+ "vision_model.encoder.layers.9.attn.qkv.bias": "native-00004.safetensors",
680
+ "vision_model.encoder.layers.9.attn.qkv.weight": "native-00004.safetensors",
681
+ "vision_model.encoder.layers.9.ls1": "native-00004.safetensors",
682
+ "vision_model.encoder.layers.9.ls2": "native-00004.safetensors",
683
+ "vision_model.encoder.layers.9.mlp.fc1.bias": "native-00004.safetensors",
684
+ "vision_model.encoder.layers.9.mlp.fc1.weight": "native-00004.safetensors",
685
+ "vision_model.encoder.layers.9.mlp.fc2.bias": "native-00004.safetensors",
686
+ "vision_model.encoder.layers.9.mlp.fc2.weight": "native-00004.safetensors",
687
+ "vision_model.encoder.layers.9.norm1.bias": "native-00004.safetensors",
688
+ "vision_model.encoder.layers.9.norm1.weight": "native-00004.safetensors",
689
+ "vision_model.encoder.layers.9.norm2.bias": "native-00004.safetensors",
690
+ "vision_model.encoder.layers.9.norm2.weight": "native-00004.safetensors"
691
+ }
692
+ }
native_backbone/modeling_intern_vit.py ADDED
@@ -0,0 +1,429 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # --------------------------------------------------------
2
+ # InternVL
3
+ # Copyright (c) 2024 OpenGVLab
4
+ # Licensed under The MIT License [see LICENSE for details]
5
+ # --------------------------------------------------------
6
+ from typing import Optional, Tuple, Union
7
+
8
+ import torch
9
+ import torch.nn.functional as F
10
+ import torch.utils.checkpoint
11
+ from einops import rearrange
12
+ from timm.layers import DropPath
13
+ from torch import nn
14
+ from transformers.activations import ACT2FN
15
+ from transformers.modeling_outputs import (BaseModelOutput,
16
+ BaseModelOutputWithPooling)
17
+ from transformers.modeling_utils import PreTrainedModel
18
+ from transformers.utils import logging
19
+
20
+ from .configuration_intern_vit import InternVisionConfig
21
+
22
+ try:
23
+ from flash_attn.bert_padding import pad_input, unpad_input
24
+ from flash_attn.flash_attn_interface import \
25
+ flash_attn_varlen_qkvpacked_func
26
+ has_flash_attn = True
27
+ except:
28
+ print('FlashAttention2 is not installed.')
29
+ has_flash_attn = False
30
+
31
+ logger = logging.get_logger(__name__)
32
+
33
+
34
+ class FlashAttention(nn.Module):
35
+ """Implement the scaled dot product attention with softmax.
36
+ Arguments
37
+ ---------
38
+ softmax_scale: The temperature to use for the softmax attention.
39
+ (default: 1/sqrt(d_keys) where d_keys is computed at
40
+ runtime)
41
+ attention_dropout: The dropout rate to apply to the attention
42
+ (default: 0.0)
43
+ """
44
+
45
+ def __init__(self, softmax_scale=None, attention_dropout=0.0, device=None, dtype=None):
46
+ super().__init__()
47
+ self.softmax_scale = softmax_scale
48
+ self.dropout_p = attention_dropout
49
+
50
+ def forward(self, qkv, key_padding_mask=None, causal=False, cu_seqlens=None,
51
+ max_s=None, need_weights=False):
52
+ """Implements the multihead softmax attention.
53
+ Arguments
54
+ ---------
55
+ qkv: The tensor containing the query, key, and value. (B, S, 3, H, D) if key_padding_mask is None
56
+ if unpadded: (nnz, 3, h, d)
57
+ key_padding_mask: a bool tensor of shape (B, S)
58
+ """
59
+ assert not need_weights
60
+ assert qkv.dtype in [torch.float16, torch.bfloat16]
61
+ assert qkv.is_cuda
62
+
63
+ if cu_seqlens is None:
64
+ batch_size = qkv.shape[0]
65
+ seqlen = qkv.shape[1]
66
+ if key_padding_mask is None:
67
+ qkv = rearrange(qkv, 'b s ... -> (b s) ...')
68
+ max_s = seqlen
69
+ cu_seqlens = torch.arange(0, (batch_size + 1) * seqlen, step=seqlen, dtype=torch.int32,
70
+ device=qkv.device)
71
+ output = flash_attn_varlen_qkvpacked_func(
72
+ qkv, cu_seqlens, max_s, self.dropout_p if self.training else 0.0,
73
+ softmax_scale=self.softmax_scale, causal=causal
74
+ )
75
+ output = rearrange(output, '(b s) ... -> b s ...', b=batch_size)
76
+ else:
77
+ nheads = qkv.shape[-2]
78
+ x = rearrange(qkv, 'b s three h d -> b s (three h d)')
79
+ x_unpad, indices, cu_seqlens, max_s = unpad_input(x, key_padding_mask)
80
+ x_unpad = rearrange(x_unpad, 'nnz (three h d) -> nnz three h d', three=3, h=nheads)
81
+ output_unpad = flash_attn_varlen_qkvpacked_func(
82
+ x_unpad, cu_seqlens, max_s, self.dropout_p if self.training else 0.0,
83
+ softmax_scale=self.softmax_scale, causal=causal
84
+ )
85
+ output = rearrange(pad_input(rearrange(output_unpad, 'nnz h d -> nnz (h d)'),
86
+ indices, batch_size, seqlen),
87
+ 'b s (h d) -> b s h d', h=nheads)
88
+ else:
89
+ assert max_s is not None
90
+ output = flash_attn_varlen_qkvpacked_func(
91
+ qkv, cu_seqlens, max_s, self.dropout_p if self.training else 0.0,
92
+ softmax_scale=self.softmax_scale, causal=causal
93
+ )
94
+
95
+ return output, None
96
+
97
+
98
+ class InternRMSNorm(nn.Module):
99
+ def __init__(self, hidden_size, eps=1e-6):
100
+ super().__init__()
101
+ self.weight = nn.Parameter(torch.ones(hidden_size))
102
+ self.variance_epsilon = eps
103
+
104
+ def forward(self, hidden_states):
105
+ input_dtype = hidden_states.dtype
106
+ hidden_states = hidden_states.to(torch.float32)
107
+ variance = hidden_states.pow(2).mean(-1, keepdim=True)
108
+ hidden_states = hidden_states * torch.rsqrt(variance + self.variance_epsilon)
109
+ return self.weight * hidden_states.to(input_dtype)
110
+
111
+
112
+ try:
113
+ from apex.normalization import FusedRMSNorm
114
+
115
+ InternRMSNorm = FusedRMSNorm # noqa
116
+
117
+ logger.info('Discovered apex.normalization.FusedRMSNorm - will use it instead of InternRMSNorm')
118
+ except ImportError:
119
+ # using the normal InternRMSNorm
120
+ pass
121
+ except Exception:
122
+ logger.warning('discovered apex but it failed to load, falling back to InternRMSNorm')
123
+ pass
124
+
125
+
126
+ NORM2FN = {
127
+ 'rms_norm': InternRMSNorm,
128
+ 'layer_norm': nn.LayerNorm,
129
+ }
130
+
131
+
132
+ class InternVisionEmbeddings(nn.Module):
133
+ def __init__(self, config: InternVisionConfig):
134
+ super().__init__()
135
+ self.config = config
136
+ self.embed_dim = config.hidden_size
137
+ self.image_size = config.image_size
138
+ self.patch_size = config.patch_size
139
+
140
+ self.class_embedding = nn.Parameter(
141
+ torch.randn(1, 1, self.embed_dim),
142
+ )
143
+
144
+ self.patch_embedding = nn.Conv2d(
145
+ in_channels=3, out_channels=self.embed_dim, kernel_size=self.patch_size, stride=self.patch_size
146
+ )
147
+
148
+ self.num_patches = (self.image_size // self.patch_size) ** 2
149
+ self.num_positions = self.num_patches + 1
150
+
151
+ self.position_embedding = nn.Parameter(torch.randn(1, self.num_positions, self.embed_dim))
152
+
153
+ def _get_pos_embed(self, pos_embed, H, W):
154
+ target_dtype = pos_embed.dtype
155
+ pos_embed = pos_embed.float().reshape(
156
+ 1, self.image_size // self.patch_size, self.image_size // self.patch_size, -1).permute(0, 3, 1, 2)
157
+ pos_embed = F.interpolate(pos_embed, size=(H, W), mode='bicubic', align_corners=False). \
158
+ reshape(1, -1, H * W).permute(0, 2, 1).to(target_dtype)
159
+ return pos_embed
160
+
161
+ def forward(self, pixel_values: torch.FloatTensor) -> torch.Tensor:
162
+ target_dtype = self.patch_embedding.weight.dtype
163
+ patch_embeds = self.patch_embedding(pixel_values) # shape = [*, channel, width, height]
164
+ batch_size, _, height, width = patch_embeds.shape
165
+ patch_embeds = patch_embeds.flatten(2).transpose(1, 2)
166
+ class_embeds = self.class_embedding.expand(batch_size, 1, -1).to(target_dtype)
167
+ embeddings = torch.cat([class_embeds, patch_embeds], dim=1)
168
+ position_embedding = torch.cat([
169
+ self.position_embedding[:, :1, :],
170
+ self._get_pos_embed(self.position_embedding[:, 1:, :], height, width)
171
+ ], dim=1)
172
+ embeddings = embeddings + position_embedding.to(target_dtype)
173
+ return embeddings
174
+
175
+
176
+ class InternAttention(nn.Module):
177
+ """Multi-headed attention from 'Attention Is All You Need' paper"""
178
+
179
+ def __init__(self, config: InternVisionConfig):
180
+ super().__init__()
181
+ self.config = config
182
+ self.embed_dim = config.hidden_size
183
+ self.num_heads = config.num_attention_heads
184
+ self.use_flash_attn = config.use_flash_attn and has_flash_attn
185
+ if config.use_flash_attn and not has_flash_attn:
186
+ print('Warning: Flash Attention is not available, use_flash_attn is set to False.')
187
+ self.head_dim = self.embed_dim // self.num_heads
188
+ if self.head_dim * self.num_heads != self.embed_dim:
189
+ raise ValueError(
190
+ f'embed_dim must be divisible by num_heads (got `embed_dim`: {self.embed_dim} and `num_heads`:'
191
+ f' {self.num_heads}).'
192
+ )
193
+
194
+ self.scale = self.head_dim ** -0.5
195
+ self.qkv = nn.Linear(self.embed_dim, 3 * self.embed_dim, bias=config.qkv_bias)
196
+ self.attn_drop = nn.Dropout(config.attention_dropout)
197
+ self.proj_drop = nn.Dropout(config.dropout)
198
+
199
+ self.qk_normalization = config.qk_normalization
200
+
201
+ if self.qk_normalization:
202
+ self.q_norm = InternRMSNorm(self.embed_dim, eps=config.layer_norm_eps)
203
+ self.k_norm = InternRMSNorm(self.embed_dim, eps=config.layer_norm_eps)
204
+
205
+ if self.use_flash_attn:
206
+ self.inner_attn = FlashAttention(attention_dropout=config.attention_dropout)
207
+ self.proj = nn.Linear(self.embed_dim, self.embed_dim)
208
+
209
+ def _naive_attn(self, x):
210
+ B, N, C = x.shape
211
+ qkv = self.qkv(x).reshape(B, N, 3, self.num_heads, C // self.num_heads).permute(2, 0, 3, 1, 4)
212
+ q, k, v = qkv.unbind(0) # make torchscript happy (cannot use tensor as tuple)
213
+
214
+ if self.qk_normalization:
215
+ B_, H_, N_, D_ = q.shape
216
+ q = self.q_norm(q.transpose(1, 2).flatten(-2, -1)).view(B_, N_, H_, D_).transpose(1, 2)
217
+ k = self.k_norm(k.transpose(1, 2).flatten(-2, -1)).view(B_, N_, H_, D_).transpose(1, 2)
218
+
219
+ attn = ((q * self.scale) @ k.transpose(-2, -1))
220
+ attn = attn.softmax(dim=-1)
221
+ attn = self.attn_drop(attn)
222
+
223
+ x = (attn @ v).transpose(1, 2).reshape(B, N, C)
224
+ x = self.proj(x)
225
+ x = self.proj_drop(x)
226
+ return x
227
+
228
+ def _flash_attn(self, x, key_padding_mask=None, need_weights=False):
229
+ qkv = self.qkv(x)
230
+ qkv = rearrange(qkv, 'b s (three h d) -> b s three h d', three=3, h=self.num_heads)
231
+
232
+ if self.qk_normalization:
233
+ q, k, v = qkv.unbind(2)
234
+ q = self.q_norm(q.flatten(-2, -1)).view(q.shape)
235
+ k = self.k_norm(k.flatten(-2, -1)).view(k.shape)
236
+ qkv = torch.stack([q, k, v], dim=2)
237
+
238
+ context, _ = self.inner_attn(
239
+ qkv, key_padding_mask=key_padding_mask, need_weights=need_weights, causal=False
240
+ )
241
+ outs = self.proj(rearrange(context, 'b s h d -> b s (h d)'))
242
+ outs = self.proj_drop(outs)
243
+ return outs
244
+
245
+ def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
246
+ x = self._naive_attn(hidden_states) if not self.use_flash_attn else self._flash_attn(hidden_states)
247
+ return x
248
+
249
+
250
+ class InternMLP(nn.Module):
251
+ def __init__(self, config: InternVisionConfig):
252
+ super().__init__()
253
+ self.config = config
254
+ self.act = ACT2FN[config.hidden_act]
255
+ self.fc1 = nn.Linear(config.hidden_size, config.intermediate_size)
256
+ self.fc2 = nn.Linear(config.intermediate_size, config.hidden_size)
257
+
258
+ def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
259
+ hidden_states = self.fc1(hidden_states)
260
+ hidden_states = self.act(hidden_states)
261
+ hidden_states = self.fc2(hidden_states)
262
+ return hidden_states
263
+
264
+
265
+ class InternVisionEncoderLayer(nn.Module):
266
+ def __init__(self, config: InternVisionConfig, drop_path_rate: float):
267
+ super().__init__()
268
+ self.embed_dim = config.hidden_size
269
+ self.intermediate_size = config.intermediate_size
270
+ self.norm_type = config.norm_type
271
+
272
+ self.attn = InternAttention(config)
273
+ self.mlp = InternMLP(config)
274
+ self.norm1 = NORM2FN[self.norm_type](self.embed_dim, eps=config.layer_norm_eps)
275
+ self.norm2 = NORM2FN[self.norm_type](self.embed_dim, eps=config.layer_norm_eps)
276
+
277
+ self.ls1 = nn.Parameter(config.initializer_factor * torch.ones(self.embed_dim))
278
+ self.ls2 = nn.Parameter(config.initializer_factor * torch.ones(self.embed_dim))
279
+ self.drop_path1 = DropPath(drop_path_rate) if drop_path_rate > 0. else nn.Identity()
280
+ self.drop_path2 = DropPath(drop_path_rate) if drop_path_rate > 0. else nn.Identity()
281
+
282
+ def forward(
283
+ self,
284
+ hidden_states: torch.Tensor,
285
+ ) -> Tuple[torch.FloatTensor, Optional[torch.FloatTensor], Optional[Tuple[torch.FloatTensor]]]:
286
+ """
287
+ Args:
288
+ hidden_states (`Tuple[torch.FloatTensor, Optional[torch.FloatTensor]]`): input to the layer of shape `(batch, seq_len, embed_dim)`
289
+ """
290
+ hidden_states = hidden_states + self.drop_path1(self.attn(self.norm1(hidden_states).to(hidden_states.dtype)) * self.ls1)
291
+
292
+ hidden_states = hidden_states + self.drop_path2(self.mlp(self.norm2(hidden_states).to(hidden_states.dtype)) * self.ls2)
293
+
294
+ return hidden_states
295
+
296
+
297
+ class InternVisionEncoder(nn.Module):
298
+ """
299
+ Transformer encoder consisting of `config.num_hidden_layers` self attention layers. Each layer is a
300
+ [`InternEncoderLayer`].
301
+
302
+ Args:
303
+ config (`InternConfig`):
304
+ The corresponding vision configuration for the `InternEncoder`.
305
+ """
306
+
307
+ def __init__(self, config: InternVisionConfig):
308
+ super().__init__()
309
+ self.config = config
310
+ # stochastic depth decay rule
311
+ dpr = [x.item() for x in torch.linspace(0, config.drop_path_rate, config.num_hidden_layers)]
312
+ self.layers = nn.ModuleList([
313
+ InternVisionEncoderLayer(config, dpr[idx]) for idx in range(config.num_hidden_layers)])
314
+ self.gradient_checkpointing = True
315
+
316
+ def forward(
317
+ self,
318
+ inputs_embeds,
319
+ output_hidden_states: Optional[bool] = None,
320
+ return_dict: Optional[bool] = None,
321
+ ) -> Union[Tuple, BaseModelOutput]:
322
+ r"""
323
+ Args:
324
+ inputs_embeds (`torch.FloatTensor` of shape `(batch_size, sequence_length, hidden_size)`):
325
+ Embedded representation of the inputs. Should be float, not int tokens.
326
+ output_hidden_states (`bool`, *optional*):
327
+ Whether or not to return the hidden states of all layers. See `hidden_states` under returned tensors
328
+ for more detail.
329
+ return_dict (`bool`, *optional*):
330
+ Whether or not to return a [`~utils.ModelOutput`] instead of a plain tuple.
331
+ """
332
+ output_hidden_states = (
333
+ output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
334
+ )
335
+ return_dict = return_dict if return_dict is not None else self.config.use_return_dict
336
+
337
+ encoder_states = () if output_hidden_states else None
338
+ hidden_states = inputs_embeds
339
+
340
+ for idx, encoder_layer in enumerate(self.layers):
341
+ if output_hidden_states:
342
+ encoder_states = encoder_states + (hidden_states,)
343
+ if self.gradient_checkpointing and self.training:
344
+ layer_outputs = torch.utils.checkpoint.checkpoint(
345
+ encoder_layer,
346
+ hidden_states)
347
+ else:
348
+ layer_outputs = encoder_layer(
349
+ hidden_states,
350
+ )
351
+ hidden_states = layer_outputs
352
+
353
+ if output_hidden_states:
354
+ encoder_states = encoder_states + (hidden_states,)
355
+
356
+ if not return_dict:
357
+ return tuple(v for v in [hidden_states, encoder_states] if v is not None)
358
+ return BaseModelOutput(
359
+ last_hidden_state=hidden_states, hidden_states=encoder_states
360
+ )
361
+
362
+
363
+ class InternVisionModel(PreTrainedModel):
364
+ main_input_name = 'pixel_values'
365
+ _supports_flash_attn_2 = True
366
+ config_class = InternVisionConfig
367
+ _no_split_modules = ['InternVisionEncoderLayer']
368
+
369
+ def __init__(self, config: InternVisionConfig):
370
+ super().__init__(config)
371
+ self.config = config
372
+
373
+ self.embeddings = InternVisionEmbeddings(config)
374
+ self.encoder = InternVisionEncoder(config)
375
+
376
+ def resize_pos_embeddings(self, old_size, new_size, patch_size):
377
+ pos_emb = self.embeddings.position_embedding
378
+ _, num_positions, embed_dim = pos_emb.shape
379
+ cls_emb = pos_emb[:, :1, :]
380
+ pos_emb = pos_emb[:, 1:, :].reshape(1, old_size // patch_size, old_size // patch_size, -1).permute(0, 3, 1, 2)
381
+ pos_emb = F.interpolate(pos_emb.float(), size=new_size // patch_size, mode='bicubic', align_corners=False)
382
+ pos_emb = pos_emb.to(cls_emb.dtype).reshape(1, embed_dim, -1).permute(0, 2, 1)
383
+ pos_emb = torch.cat([cls_emb, pos_emb], dim=1)
384
+ self.embeddings.position_embedding = nn.Parameter(pos_emb)
385
+ self.embeddings.image_size = new_size
386
+ logger.info('Resized position embeddings from {} to {}'.format(old_size, new_size))
387
+
388
+ def get_input_embeddings(self):
389
+ return self.embeddings
390
+
391
+ def forward(
392
+ self,
393
+ pixel_values: Optional[torch.FloatTensor] = None,
394
+ output_hidden_states: Optional[bool] = None,
395
+ return_dict: Optional[bool] = None,
396
+ pixel_embeds: Optional[torch.FloatTensor] = None,
397
+ ) -> Union[Tuple, BaseModelOutputWithPooling]:
398
+ output_hidden_states = (
399
+ output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
400
+ )
401
+ return_dict = return_dict if return_dict is not None else self.config.use_return_dict
402
+
403
+ if pixel_values is None and pixel_embeds is None:
404
+ raise ValueError('You have to specify pixel_values or pixel_embeds')
405
+
406
+ if pixel_embeds is not None:
407
+ hidden_states = pixel_embeds
408
+ else:
409
+ if len(pixel_values.shape) == 4:
410
+ hidden_states = self.embeddings(pixel_values)
411
+ else:
412
+ raise ValueError(f'wrong pixel_values size: {pixel_values.shape}')
413
+ encoder_outputs = self.encoder(
414
+ inputs_embeds=hidden_states,
415
+ output_hidden_states=output_hidden_states,
416
+ return_dict=return_dict,
417
+ )
418
+ last_hidden_state = encoder_outputs.last_hidden_state
419
+ pooled_output = last_hidden_state[:, 0, :]
420
+
421
+ if not return_dict:
422
+ return (last_hidden_state, pooled_output) + encoder_outputs[1:]
423
+
424
+ return BaseModelOutputWithPooling(
425
+ last_hidden_state=last_hidden_state,
426
+ pooler_output=pooled_output,
427
+ hidden_states=encoder_outputs.hidden_states,
428
+ attentions=encoder_outputs.attentions,
429
+ )
native_backbone/modeling_internlm2.py ADDED
@@ -0,0 +1,1456 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) The InternLM team and The HuggingFace Inc. team. All rights reserved.
2
+ #
3
+ # This code is based on transformers/src/transformers/models/llama/modeling_llama.py
4
+ #
5
+ # Licensed under the Apache License, Version 2.0 (the "License");
6
+ # you may not use this file except in compliance with the License.
7
+ # You may obtain a copy of the License at
8
+ #
9
+ # http://www.apache.org/licenses/LICENSE-2.0
10
+ #
11
+ # Unless required by applicable law or agreed to in writing, software
12
+ # distributed under the License is distributed on an "AS IS" BASIS,
13
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
14
+ # See the License for the specific language governing permissions and
15
+ # limitations under the License.
16
+ """ PyTorch InternLM2 model."""
17
+ import math
18
+ import queue
19
+ import threading
20
+ import warnings
21
+ from typing import List, Optional, Tuple, Union
22
+
23
+ import torch
24
+ import torch.nn.functional as F
25
+ import torch.utils.checkpoint
26
+ from einops import rearrange
27
+ from torch import nn
28
+ from torch.nn import BCEWithLogitsLoss, CrossEntropyLoss, MSELoss
29
+ from transformers.activations import ACT2FN
30
+ from transformers.modeling_outputs import (BaseModelOutputWithPast,
31
+ CausalLMOutputWithPast,
32
+ SequenceClassifierOutputWithPast)
33
+ from transformers.modeling_utils import PreTrainedModel
34
+ from transformers.utils import (add_start_docstrings,
35
+ add_start_docstrings_to_model_forward, logging,
36
+ replace_return_docstrings)
37
+
38
+ try:
39
+ from transformers.generation.streamers import BaseStreamer
40
+ except: # noqa # pylint: disable=bare-except
41
+ BaseStreamer = None
42
+
43
+ from .configuration_internlm2 import InternLM2Config
44
+
45
+ logger = logging.get_logger(__name__)
46
+
47
+ _CONFIG_FOR_DOC = 'InternLM2Config'
48
+
49
+ flash_attn_func, flash_attn_varlen_func = None, None
50
+ pad_input, index_first_axis, unpad_input = None, None, None
51
+ try:
52
+ from flash_attn import flash_attn_func as _flash_attn_func
53
+ from flash_attn import flash_attn_varlen_func as _flash_attn_varlen_func
54
+ from flash_attn.bert_padding import index_first_axis as _index_first_axis
55
+ from flash_attn.bert_padding import pad_input as _pad_input
56
+ from flash_attn.bert_padding import unpad_input as _unpad_input
57
+
58
+ flash_attn_func, flash_attn_varlen_func = _flash_attn_func, _flash_attn_varlen_func
59
+ pad_input, index_first_axis, unpad_input = _pad_input, _index_first_axis, _unpad_input
60
+ has_flash_attn = True
61
+ except:
62
+ has_flash_attn = False
63
+
64
+
65
+ def _import_flash_attn():
66
+ global flash_attn_func, flash_attn_varlen_func
67
+ global pad_input, index_first_axis, unpad_input
68
+ try:
69
+ from flash_attn import flash_attn_func as _flash_attn_func
70
+ from flash_attn import \
71
+ flash_attn_varlen_func as _flash_attn_varlen_func
72
+ from flash_attn.bert_padding import \
73
+ index_first_axis as _index_first_axis
74
+ from flash_attn.bert_padding import pad_input as _pad_input
75
+ from flash_attn.bert_padding import unpad_input as _unpad_input
76
+ flash_attn_func, flash_attn_varlen_func = _flash_attn_func, _flash_attn_varlen_func
77
+ pad_input, index_first_axis, unpad_input = _pad_input, _index_first_axis, _unpad_input
78
+ except ImportError:
79
+ raise ImportError('flash_attn is not installed.')
80
+
81
+
82
+ # Copied from transformers.models.llama.modeling_llama._get_unpad_data
83
+ def _get_unpad_data(attention_mask):
84
+ seqlens_in_batch = attention_mask.sum(dim=-1, dtype=torch.int32)
85
+ indices = torch.nonzero(attention_mask.flatten(), as_tuple=False).flatten()
86
+ max_seqlen_in_batch = seqlens_in_batch.max().item()
87
+ cu_seqlens = F.pad(torch.cumsum(seqlens_in_batch, dim=0, dtype=torch.torch.int32), (1, 0))
88
+ return (
89
+ indices,
90
+ cu_seqlens,
91
+ max_seqlen_in_batch,
92
+ )
93
+
94
+
95
+ # Copied from transformers.models.bart.modeling_bart._make_causal_mask
96
+ def _make_causal_mask(
97
+ input_ids_shape: torch.Size, dtype: torch.dtype, device: torch.device, past_key_values_length: int = 0
98
+ ):
99
+ """
100
+ Make causal mask used for bi-directional self-attention.
101
+ """
102
+ bsz, tgt_len = input_ids_shape
103
+ mask = torch.full((tgt_len, tgt_len), torch.tensor(torch.finfo(dtype).min, device=device), device=device)
104
+ mask_cond = torch.arange(mask.size(-1), device=device)
105
+ mask.masked_fill_(mask_cond < (mask_cond + 1).view(mask.size(-1), 1), 0)
106
+ mask = mask.to(dtype)
107
+
108
+ if past_key_values_length > 0:
109
+ mask = torch.cat([torch.zeros(tgt_len, past_key_values_length, dtype=dtype, device=device), mask], dim=-1)
110
+ return mask[None, None, :, :].expand(bsz, 1, tgt_len, tgt_len + past_key_values_length)
111
+
112
+
113
+ # Copied from transformers.models.bart.modeling_bart._expand_mask
114
+ def _expand_mask(mask: torch.Tensor, dtype: torch.dtype, tgt_len: Optional[int] = None):
115
+ """
116
+ Expands attention_mask from `[bsz, seq_len]` to `[bsz, 1, tgt_seq_len, src_seq_len]`.
117
+ """
118
+ bsz, src_len = mask.size()
119
+ tgt_len = tgt_len if tgt_len is not None else src_len
120
+
121
+ expanded_mask = mask[:, None, None, :].expand(bsz, 1, tgt_len, src_len).to(dtype)
122
+
123
+ inverted_mask = 1.0 - expanded_mask
124
+
125
+ return inverted_mask.masked_fill(inverted_mask.to(torch.bool), torch.finfo(dtype).min)
126
+
127
+
128
+ # Copied from transformers.models.llama.modeling_llama.LlamaRMSNorm with Llama->InternLM2
129
+ class InternLM2RMSNorm(nn.Module):
130
+ def __init__(self, hidden_size, eps=1e-6):
131
+ """
132
+ InternLM2RMSNorm is equivalent to T5LayerNorm
133
+ """
134
+ super().__init__()
135
+ self.weight = nn.Parameter(torch.ones(hidden_size))
136
+ self.variance_epsilon = eps
137
+
138
+ def forward(self, hidden_states):
139
+ input_dtype = hidden_states.dtype
140
+ hidden_states = hidden_states.to(torch.float32)
141
+ variance = hidden_states.pow(2).mean(-1, keepdim=True)
142
+ hidden_states = hidden_states * torch.rsqrt(variance + self.variance_epsilon)
143
+ return self.weight * hidden_states.to(input_dtype)
144
+
145
+
146
+ try:
147
+ from functools import partial
148
+
149
+ from apex.normalization import FusedRMSNorm
150
+ InternLM2RMSNorm = partial(FusedRMSNorm, eps=1e-6) # noqa
151
+ print('Discovered apex.normalization.FusedRMSNorm - will use it instead of InternLM2RMSNorm')
152
+ except ImportError:
153
+ # using the normal LlamaRMSNorm
154
+ pass
155
+ except Exception:
156
+ print('discovered apex but it failed to load, falling back to InternLM2RMSNorm')
157
+ pass
158
+
159
+
160
+ # Copied from transformers.model.llama.modeling_llama.LlamaRotaryEmbedding with Llama->InternLM2
161
+ class InternLM2RotaryEmbedding(nn.Module):
162
+ def __init__(self, dim, max_position_embeddings=2048, base=10000, device=None):
163
+ super().__init__()
164
+
165
+ self.dim = dim
166
+ self.max_position_embeddings = max_position_embeddings
167
+ self.base = base
168
+ self.inv_freq = None
169
+ # inv_freq = 1.0 / (self.base ** (torch.arange(0, self.dim, 2).float().to(device) / self.dim))
170
+ # self.register_buffer('inv_freq', inv_freq, persistent=False)
171
+
172
+ self.max_seq_len_cached = -1
173
+ # Build here to make `torch.jit.trace` work.
174
+ # self._set_cos_sin_cache(
175
+ # seq_len=max_position_embeddings, device=self.inv_freq.device, dtype=torch.get_default_dtype()
176
+ # )
177
+
178
+ def _set_cos_sin_cache(self, seq_len, device, dtype):
179
+ if self.inv_freq is None:
180
+ inv_freq = 1.0 / (self.base ** (torch.arange(0, self.dim, 2, device=device, dtype=torch.float32) / self.dim))
181
+ del self.inv_freq
182
+ self.register_buffer('inv_freq', inv_freq, persistent=False)
183
+
184
+
185
+ self.max_seq_len_cached = seq_len
186
+ t = torch.arange(self.max_seq_len_cached, device=device).to(dtype=self.inv_freq.dtype)
187
+
188
+ # freqs = torch.einsum('i,j->ij', t, self.inv_freq)
189
+ freqs = torch.outer(t, self.inv_freq.to(device=t.device))
190
+
191
+ # Different from paper, but it uses a different permutation in order to obtain the same calculation
192
+ emb = torch.cat((freqs, freqs), dim=-1)
193
+ self.register_buffer('cos_cached', emb.cos().to(dtype), persistent=False)
194
+ self.register_buffer('sin_cached', emb.sin().to(dtype), persistent=False)
195
+
196
+ def forward(self, x, seq_len=None):
197
+ # x: [bs, num_attention_heads, seq_len, head_size]
198
+ if seq_len > self.max_seq_len_cached:
199
+ self._set_cos_sin_cache(seq_len=seq_len, device=x.device, dtype=x.dtype)
200
+
201
+ return (
202
+ self.cos_cached[:seq_len].to(dtype=x.dtype),
203
+ self.sin_cached[:seq_len].to(dtype=x.dtype),
204
+ )
205
+
206
+
207
+ # Copied from transformers.model.llama.modeling_llama.LlamaLinearScalingRotaryEmbedding with Llama->InternLM2
208
+ class InternLM2LinearScalingRotaryEmbedding(InternLM2RotaryEmbedding):
209
+ """InternLM2RotaryEmbedding extended with linear scaling. Credits to the Reddit user /u/kaiokendev"""
210
+
211
+ def __init__(self, dim, max_position_embeddings=2048, base=10000, device=None, scaling_factor=1.0):
212
+ self.scaling_factor = scaling_factor
213
+ super().__init__(dim, max_position_embeddings, base, device)
214
+
215
+ def _set_cos_sin_cache(self, seq_len, device, dtype):
216
+ if self.inv_freq is None:
217
+ inv_freq = 1.0 / (self.base ** (torch.arange(0, self.dim, 2, device=device, dtype=torch.float32) / self.dim))
218
+ del self.inv_freq
219
+ self.register_buffer('inv_freq', inv_freq, persistent=False)
220
+
221
+ self.max_seq_len_cached = seq_len
222
+ t = torch.arange(self.max_seq_len_cached, device=device).to(dtype=self.inv_freq.dtype)
223
+ t = t / self.scaling_factor
224
+
225
+ # freqs = torch.einsum('i,j->ij', t, self.inv_freq)
226
+ freqs = torch.outer(t, self.inv_freq.to(device=t.device))
227
+
228
+ # Different from paper, but it uses a different permutation in order to obtain the same calculation
229
+ emb = torch.cat((freqs, freqs), dim=-1)
230
+ self.register_buffer('cos_cached', emb.cos().to(dtype), persistent=False)
231
+ self.register_buffer('sin_cached', emb.sin().to(dtype), persistent=False)
232
+
233
+
234
+ # Copied from transformers.model.llama.modeling_llama.LlamaDynamicNTKScalingRotaryEmbedding with Llama->InternLM2
235
+ class InternLM2DynamicNTKScalingRotaryEmbedding(InternLM2RotaryEmbedding):
236
+ """InternLM2RotaryEmbedding extended with Dynamic NTK scaling.
237
+ Credits to the Reddit users /u/bloc97 and /u/emozilla.
238
+ """
239
+
240
+ def __init__(self, dim, max_position_embeddings=2048, base=10000, device=None, scaling_factor=1.0):
241
+ self.scaling_factor = scaling_factor
242
+ super().__init__(dim, max_position_embeddings, base, device)
243
+
244
+ def _set_cos_sin_cache(self, seq_len, device, dtype):
245
+ if self.inv_freq is None:
246
+ inv_freq = 1.0 / (self.base ** (torch.arange(0, self.dim, 2, device=device, dtype=torch.float32) / self.dim))
247
+ del self.inv_freq
248
+ self.register_buffer('inv_freq', inv_freq, persistent=False)
249
+
250
+
251
+ self.max_seq_len_cached = seq_len
252
+
253
+ if seq_len > self.max_position_embeddings:
254
+ base = self.base * (
255
+ (self.scaling_factor * seq_len / self.max_position_embeddings) - (self.scaling_factor - 1)
256
+ ) ** (self.dim / (self.dim - 2))
257
+ inv_freq = 1.0 / (base ** (torch.arange(0, self.dim, 2).float().to(device) / self.dim))
258
+ self.register_buffer('inv_freq', inv_freq, persistent=False)
259
+
260
+ t = torch.arange(self.max_seq_len_cached, device=device).to(dtype=self.inv_freq.dtype)
261
+
262
+ # freqs = torch.einsum('i,j->ij', t, self.inv_freq)
263
+ freqs = torch.outer(t, self.inv_freq.to(device=t.device))
264
+
265
+ # Different from paper, but it uses a different permutation in order to obtain the same calculation
266
+ emb = torch.cat((freqs, freqs), dim=-1)
267
+ self.register_buffer('cos_cached', emb.cos().to(dtype), persistent=False)
268
+ self.register_buffer('sin_cached', emb.sin().to(dtype), persistent=False)
269
+
270
+
271
+ # Copied from transformers.model.llama.modeling_llama.rotate_half
272
+ def rotate_half(x):
273
+ """Rotates half the hidden dims of the input."""
274
+ x1 = x[..., : x.shape[-1] // 2]
275
+ x2 = x[..., x.shape[-1] // 2:]
276
+ return torch.cat((-x2, x1), dim=-1)
277
+
278
+
279
+
280
+ # Copied from transformers.model.llama.modeling_llama.apply_rotary_pos_emb; float
281
+ def apply_rotary_pos_emb(q, k, cos, sin, position_ids, unsqueeze_dim=1):
282
+ """Applies Rotary Position Embedding to the query and key tensors."""
283
+ cos = cos[position_ids].unsqueeze(unsqueeze_dim).float()
284
+ sin = sin[position_ids].unsqueeze(unsqueeze_dim).float()
285
+ q_dtype, k_dtype = q.dtype, k.dtype
286
+ q, k = q.float(), k.float()
287
+ q_embed = (q * cos) + (rotate_half(q) * sin)
288
+ k_embed = (k * cos) + (rotate_half(k) * sin)
289
+ return q_embed.to(dtype=q_dtype), k_embed.to(dtype=k_dtype)
290
+
291
+
292
+ class InternLM2MLP(nn.Module):
293
+ def __init__(self, config):
294
+ super().__init__()
295
+ self.config = config
296
+ self.hidden_size = config.hidden_size
297
+ self.intermediate_size = config.intermediate_size
298
+ self.w1 = nn.Linear(self.hidden_size, self.intermediate_size, bias=False)
299
+ self.w3 = nn.Linear(self.hidden_size, self.intermediate_size, bias=False)
300
+ self.w2 = nn.Linear(self.intermediate_size, self.hidden_size, bias=False)
301
+ self.act_fn = ACT2FN[config.hidden_act]
302
+
303
+ def forward(self, x):
304
+ down_proj = self.w2(self.act_fn(self.w1(x)) * self.w3(x))
305
+
306
+ return down_proj
307
+
308
+
309
+ # Copied from transformers.model.llama.modeling_llama.repeat_kv
310
+ def repeat_kv(hidden_states: torch.Tensor, n_rep: int) -> torch.Tensor:
311
+ """
312
+ This is the equivalent of torch.repeat_interleave(x, dim=1, repeats=n_rep). The hidden states go from (batch,
313
+ num_key_value_heads, seqlen, head_dim) to (batch, num_attention_heads, seqlen, head_dim)
314
+ """
315
+ batch, num_key_value_heads, slen, head_dim = hidden_states.shape
316
+ if n_rep == 1:
317
+ return hidden_states
318
+ hidden_states = hidden_states[:, :, None, :, :].expand(batch, num_key_value_heads, n_rep, slen, head_dim)
319
+ return hidden_states.reshape(batch, num_key_value_heads * n_rep, slen, head_dim)
320
+
321
+
322
+ # Modified from transformers.model.llama.modeling_llama.LlamaAttention
323
+ class InternLM2Attention(nn.Module):
324
+ """Multi-headed attention from 'Attention Is All You Need' paper"""
325
+
326
+ def __init__(self, config: InternLM2Config):
327
+ super().__init__()
328
+ self.config = config
329
+ self.hidden_size = config.hidden_size
330
+ self.num_heads = config.num_attention_heads
331
+ self.head_dim = self.hidden_size // self.num_heads
332
+ self.num_key_value_heads = config.num_key_value_heads
333
+ self.num_key_value_groups = self.num_heads // self.num_key_value_heads
334
+ self.max_position_embeddings = config.max_position_embeddings
335
+ self.is_causal = True
336
+
337
+ if (self.head_dim * self.num_heads) != self.hidden_size:
338
+ raise ValueError(
339
+ f'hidden_size must be divisible by num_heads (got `hidden_size`: {self.hidden_size}'
340
+ f' and `num_heads`: {self.num_heads}).'
341
+ )
342
+
343
+ self.wqkv = nn.Linear(
344
+ self.hidden_size,
345
+ (self.num_heads + 2 * self.num_key_value_heads) * self.head_dim,
346
+ bias=config.bias,
347
+ )
348
+
349
+ self.wo = nn.Linear(self.num_heads * self.head_dim, self.hidden_size, bias=config.bias)
350
+ self._init_rope()
351
+
352
+ def _init_rope(self):
353
+ if self.config.rope_scaling is None:
354
+ self.rotary_emb = InternLM2RotaryEmbedding(
355
+ self.head_dim,
356
+ max_position_embeddings=self.max_position_embeddings,
357
+ base=self.config.rope_theta,
358
+ )
359
+ else:
360
+ scaling_type = self.config.rope_scaling['type']
361
+ scaling_factor = self.config.rope_scaling['factor']
362
+ if scaling_type == 'dynamic':
363
+ self.rotary_emb = InternLM2DynamicNTKScalingRotaryEmbedding(
364
+ self.head_dim,
365
+ max_position_embeddings=self.max_position_embeddings,
366
+ base=self.config.rope_theta,
367
+ scaling_factor=scaling_factor,
368
+ )
369
+ elif scaling_type == 'linear':
370
+ self.rotary_emb = InternLM2LinearScalingRotaryEmbedding(
371
+ self.head_dim,
372
+ max_position_embeddings=self.max_position_embeddings,
373
+ base=self.config.rope_theta,
374
+ scaling_factor=scaling_factor,
375
+ )
376
+ else:
377
+ raise ValueError("Currently we only support rotary embedding's type being 'dynamic' or 'linear'.")
378
+ return self.rotary_emb
379
+
380
+ def _shape(self, tensor: torch.Tensor, seq_len: int, bsz: int):
381
+ return tensor.view(bsz, seq_len, self.num_heads, self.head_dim).transpose(1, 2).contiguous()
382
+
383
+ def forward(
384
+ self,
385
+ hidden_states: torch.Tensor,
386
+ attention_mask: Optional[torch.Tensor] = None,
387
+ position_ids: Optional[torch.LongTensor] = None,
388
+ past_key_value: Optional[Tuple[torch.Tensor]] = None,
389
+ output_attentions: bool = False,
390
+ use_cache: bool = False,
391
+ **kwargs,
392
+ ) -> Tuple[torch.Tensor, Optional[torch.Tensor], Optional[Tuple[torch.Tensor]]]:
393
+ if 'padding_mask' in kwargs:
394
+ warnings.warn(
395
+ 'Passing `padding_mask` is deprecated and will be removed in v4.37. '
396
+ 'Please make sure use `attention_mask` instead.`'
397
+ )
398
+
399
+ bsz, q_len, _ = hidden_states.size()
400
+
401
+ qkv_states = self.wqkv(hidden_states)
402
+
403
+ qkv_states = rearrange(
404
+ qkv_states,
405
+ 'b q (h gs d) -> b q h gs d',
406
+ gs=2 + self.num_key_value_groups,
407
+ d=self.head_dim,
408
+ )
409
+
410
+ query_states = qkv_states[..., : self.num_key_value_groups, :]
411
+ query_states = rearrange(query_states, 'b q h gs d -> b q (h gs) d')
412
+ key_states = qkv_states[..., -2, :]
413
+ value_states = qkv_states[..., -1, :]
414
+
415
+ query_states = query_states.transpose(1, 2)
416
+ key_states = key_states.transpose(1, 2)
417
+ value_states = value_states.transpose(1, 2)
418
+
419
+ kv_seq_len = key_states.shape[-2]
420
+ if past_key_value is not None:
421
+ kv_seq_len += past_key_value[0].shape[-2]
422
+ cos, sin = self.rotary_emb(value_states, seq_len=kv_seq_len)
423
+ query_states, key_states = apply_rotary_pos_emb(query_states, key_states, cos, sin, position_ids)
424
+
425
+ if past_key_value is not None:
426
+ # reuse k, v, self_attention
427
+ key_states = torch.cat([past_key_value[0], key_states], dim=2)
428
+ value_states = torch.cat([past_key_value[1], value_states], dim=2)
429
+
430
+ past_key_value = (key_states, value_states) if use_cache else None
431
+
432
+ key_states = repeat_kv(key_states, self.num_key_value_groups)
433
+ value_states = repeat_kv(value_states, self.num_key_value_groups)
434
+
435
+ attn_weights = torch.matmul(query_states, key_states.transpose(2, 3)) / math.sqrt(self.head_dim)
436
+
437
+ if attn_weights.size() != (bsz, self.num_heads, q_len, kv_seq_len):
438
+ raise ValueError(
439
+ f'Attention weights should be of size {(bsz, self.num_heads, q_len, kv_seq_len)}, but is'
440
+ f' {attn_weights.size()}'
441
+ )
442
+
443
+ if attention_mask is not None:
444
+ if attention_mask.size() != (bsz, 1, q_len, kv_seq_len):
445
+ raise ValueError(
446
+ f'Attention mask should be of size {(bsz, 1, q_len, kv_seq_len)}, but is {attention_mask.size()}'
447
+ )
448
+ attn_weights = attn_weights + attention_mask
449
+
450
+ # upcast attention to fp32
451
+ attn_weights = nn.functional.softmax(attn_weights, dim=-1, dtype=torch.float32).to(query_states.dtype)
452
+ attn_output = torch.matmul(attn_weights, value_states)
453
+
454
+ if attn_output.size() != (bsz, self.num_heads, q_len, self.head_dim):
455
+ raise ValueError(
456
+ f'`attn_output` should be of size {(bsz, self.num_heads, q_len, self.head_dim)}, but is'
457
+ f' {attn_output.size()}'
458
+ )
459
+
460
+ attn_output = attn_output.transpose(1, 2).contiguous()
461
+ attn_output = attn_output.reshape(bsz, q_len, self.hidden_size)
462
+
463
+ attn_output = self.wo(attn_output)
464
+
465
+ if not output_attentions:
466
+ attn_weights = None
467
+
468
+ return attn_output, attn_weights, past_key_value
469
+
470
+
471
+ # Modified from transformers.model.llama.modeling_llama.InternLM2FlashAttention2
472
+ class InternLM2FlashAttention2(InternLM2Attention):
473
+ """
474
+ InternLM2 flash attention module. This module inherits from `InternLM2Attention` as the weights of the module stays
475
+ untouched. The only required change would be on the forward pass where it needs to correctly call the public API of
476
+ flash attention and deal with padding tokens in case the input contains any of them.
477
+ """
478
+
479
+ def forward(
480
+ self,
481
+ hidden_states: torch.Tensor,
482
+ attention_mask: Optional[torch.LongTensor] = None,
483
+ position_ids: Optional[torch.LongTensor] = None,
484
+ past_key_value: Optional[Tuple[torch.Tensor]] = None,
485
+ output_attentions: bool = False,
486
+ use_cache: bool = False,
487
+ **kwargs,
488
+ ) -> Tuple[torch.Tensor, Optional[torch.Tensor], Optional[Tuple[torch.Tensor]]]:
489
+ # InternLM2FlashAttention2 attention does not support output_attentions
490
+ if 'padding_mask' in kwargs:
491
+ warnings.warn(
492
+ 'Passing `padding_mask` is deprecated and will be removed in v4.37. '
493
+ 'Please make sure use `attention_mask` instead.`'
494
+ )
495
+
496
+ # overwrite attention_mask with padding_mask
497
+ attention_mask = kwargs.pop('padding_mask')
498
+
499
+ output_attentions = False
500
+
501
+ bsz, q_len, _ = hidden_states.size()
502
+
503
+ qkv_states = self.wqkv(hidden_states)
504
+
505
+ qkv_states = rearrange(
506
+ qkv_states,
507
+ 'b q (h gs d) -> b q h gs d',
508
+ gs=2 + self.num_key_value_groups,
509
+ d=self.head_dim,
510
+ )
511
+
512
+ query_states = qkv_states[..., : self.num_key_value_groups, :]
513
+ query_states = rearrange(query_states, 'b q h gs d -> b q (h gs) d')
514
+ key_states = qkv_states[..., -2, :]
515
+ value_states = qkv_states[..., -1, :]
516
+
517
+ query_states = query_states.transpose(1, 2)
518
+ key_states = key_states.transpose(1, 2)
519
+ value_states = value_states.transpose(1, 2)
520
+
521
+ kv_seq_len = key_states.shape[-2]
522
+ if past_key_value is not None:
523
+ kv_seq_len += past_key_value[0].shape[-2]
524
+
525
+ cos, sin = self.rotary_emb(value_states, seq_len=kv_seq_len)
526
+
527
+ query_states, key_states = apply_rotary_pos_emb(query_states, key_states, cos, sin, position_ids)
528
+
529
+ if past_key_value is not None:
530
+ # reuse k, v, self_attention
531
+ key_states = torch.cat([past_key_value[0], key_states], dim=2)
532
+ value_states = torch.cat([past_key_value[1], value_states], dim=2)
533
+
534
+ past_key_value = (key_states, value_states) if use_cache else None
535
+
536
+ query_states = query_states.transpose(1, 2)
537
+ key_states = key_states.transpose(1, 2)
538
+ value_states = value_states.transpose(1, 2)
539
+
540
+ attn_output = self._flash_attention_forward(
541
+ query_states, key_states, value_states, attention_mask, q_len
542
+ )
543
+ attn_output = attn_output.reshape(bsz, q_len, self.hidden_size).contiguous()
544
+ attn_output = self.wo(attn_output)
545
+
546
+ if not output_attentions:
547
+ attn_weights = None
548
+
549
+ return attn_output, attn_weights, past_key_value
550
+
551
+ def _flash_attention_forward(
552
+ self, query_states, key_states, value_states, attention_mask, query_length, dropout=0.0, softmax_scale=None
553
+ ):
554
+ """
555
+ Calls the forward method of Flash Attention - if the input hidden states contain at least one padding token
556
+ first unpad the input, then computes the attention scores and pad the final attention scores.
557
+
558
+ Args:
559
+ query_states (`torch.Tensor`):
560
+ Input query states to be passed to Flash Attention API
561
+ key_states (`torch.Tensor`):
562
+ Input key states to be passed to Flash Attention API
563
+ value_states (`torch.Tensor`):
564
+ Input value states to be passed to Flash Attention API
565
+ attention_mask (`torch.Tensor`):
566
+ The padding mask - corresponds to a tensor of size `(batch_size, seq_len)` where 0 stands for the
567
+ position of padding tokens and 1 for the position of non-padding tokens.
568
+ dropout (`int`, *optional*):
569
+ Attention dropout
570
+ softmax_scale (`float`, *optional*):
571
+ The scaling of QK^T before applying softmax. Default to 1 / sqrt(head_dim)
572
+ """
573
+ # Contains at least one padding token in the sequence
574
+ causal = self.is_causal and query_length != 1
575
+ if attention_mask is not None:
576
+ batch_size = query_states.shape[0]
577
+ query_states, key_states, value_states, indices_q, cu_seq_lens, max_seq_lens = self._unpad_input(
578
+ query_states, key_states, value_states, attention_mask, query_length
579
+ )
580
+
581
+ cu_seqlens_q, cu_seqlens_k = cu_seq_lens
582
+ max_seqlen_in_batch_q, max_seqlen_in_batch_k = max_seq_lens
583
+
584
+ attn_output_unpad = flash_attn_varlen_func(
585
+ query_states,
586
+ key_states,
587
+ value_states,
588
+ cu_seqlens_q=cu_seqlens_q,
589
+ cu_seqlens_k=cu_seqlens_k,
590
+ max_seqlen_q=max_seqlen_in_batch_q,
591
+ max_seqlen_k=max_seqlen_in_batch_k,
592
+ dropout_p=dropout,
593
+ softmax_scale=softmax_scale,
594
+ causal=causal,
595
+ )
596
+
597
+ attn_output = pad_input(attn_output_unpad, indices_q, batch_size, query_length)
598
+ else:
599
+ attn_output = flash_attn_func(
600
+ query_states, key_states, value_states, dropout, softmax_scale=softmax_scale, causal=causal
601
+ )
602
+
603
+ return attn_output
604
+
605
+ def _unpad_input(self, query_layer, key_layer, value_layer, attention_mask, query_length):
606
+ indices_k, cu_seqlens_k, max_seqlen_in_batch_k = _get_unpad_data(attention_mask)
607
+ batch_size, kv_seq_len, num_key_value_heads, head_dim = key_layer.shape
608
+
609
+ key_layer = index_first_axis(
610
+ key_layer.reshape(batch_size * kv_seq_len, num_key_value_heads, head_dim), indices_k
611
+ )
612
+ value_layer = index_first_axis(
613
+ value_layer.reshape(batch_size * kv_seq_len, num_key_value_heads, head_dim), indices_k
614
+ )
615
+
616
+ if query_length == kv_seq_len:
617
+ query_layer = index_first_axis(
618
+ query_layer.reshape(batch_size * kv_seq_len, self.num_heads, head_dim), indices_k
619
+ )
620
+ cu_seqlens_q = cu_seqlens_k
621
+ max_seqlen_in_batch_q = max_seqlen_in_batch_k
622
+ indices_q = indices_k
623
+ elif query_length == 1:
624
+ max_seqlen_in_batch_q = 1
625
+ cu_seqlens_q = torch.arange(
626
+ batch_size + 1, dtype=torch.int32, device=query_layer.device
627
+ ) # There is a memcpy here, that is very bad.
628
+ indices_q = cu_seqlens_q[:-1]
629
+ query_layer = query_layer.squeeze(1)
630
+ else:
631
+ # The -q_len: slice assumes left padding.
632
+ attention_mask = attention_mask[:, -query_length:]
633
+ query_layer, indices_q, cu_seqlens_q, max_seqlen_in_batch_q = unpad_input(query_layer, attention_mask)
634
+
635
+ return (
636
+ query_layer,
637
+ key_layer,
638
+ value_layer,
639
+ indices_q.to(torch.int64),
640
+ (cu_seqlens_q, cu_seqlens_k),
641
+ (max_seqlen_in_batch_q, max_seqlen_in_batch_k),
642
+ )
643
+
644
+
645
+ INTERNLM2_ATTENTION_CLASSES = {
646
+ 'eager': InternLM2Attention,
647
+ 'flash_attention_2': InternLM2FlashAttention2,
648
+ }
649
+
650
+
651
+ # Modified from transformers.model.llama.modeling_llama.LlamaDecoderLayer
652
+ class InternLM2DecoderLayer(nn.Module):
653
+ def __init__(self, config: InternLM2Config):
654
+ super().__init__()
655
+ self.hidden_size = config.hidden_size
656
+
657
+ self.attention = INTERNLM2_ATTENTION_CLASSES[config.attn_implementation](config=config)
658
+
659
+ self.feed_forward = InternLM2MLP(config)
660
+ self.attention_norm = InternLM2RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
661
+ self.ffn_norm = InternLM2RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
662
+
663
+ def forward(
664
+ self,
665
+ hidden_states: torch.Tensor,
666
+ attention_mask: Optional[torch.Tensor] = None,
667
+ position_ids: Optional[torch.LongTensor] = None,
668
+ past_key_value: Optional[Tuple[torch.Tensor]] = None,
669
+ output_attentions: Optional[bool] = False,
670
+ use_cache: Optional[bool] = False,
671
+ **kwargs,
672
+ ) -> Tuple[torch.FloatTensor, Optional[Tuple[torch.FloatTensor, torch.FloatTensor]]]:
673
+ """
674
+ Args:
675
+ hidden_states (`torch.FloatTensor`): input to the layer of shape `(batch, seq_len, embed_dim)`
676
+ attention_mask (`torch.FloatTensor`, *optional*):
677
+ attention mask of size `(batch_size, sequence_length)` if flash attention is used or `(batch_size, 1,
678
+ query_sequence_length, key_sequence_length)` if default attention is used.
679
+ output_attentions (`bool`, *optional*):
680
+ Whether or not to return the attentions tensors of all attention layers. See `attentions` under
681
+ returned tensors for more detail.
682
+ use_cache (`bool`, *optional*):
683
+ If set to `True`, `past_key_values` key value states are returned and can be used to speed up decoding
684
+ (see `past_key_values`).
685
+ past_key_value (`Tuple(torch.FloatTensor)`, *optional*): cached past key and value projection states
686
+ """
687
+ if 'padding_mask' in kwargs:
688
+ warnings.warn(
689
+ 'Passing `padding_mask` is deprecated and will be removed in v4.37. '
690
+ 'Please make sure use `attention_mask` instead.`'
691
+ )
692
+
693
+ residual = hidden_states
694
+
695
+ hidden_states = self.attention_norm(hidden_states)
696
+
697
+ # Self Attention
698
+ hidden_states, self_attn_weights, present_key_value = self.attention(
699
+ hidden_states=hidden_states,
700
+ attention_mask=attention_mask,
701
+ position_ids=position_ids,
702
+ past_key_value=past_key_value,
703
+ output_attentions=output_attentions,
704
+ use_cache=use_cache,
705
+ **kwargs,
706
+ )
707
+ hidden_states = residual + hidden_states
708
+
709
+ # Fully Connected
710
+ residual = hidden_states
711
+ hidden_states = self.ffn_norm(hidden_states)
712
+ hidden_states = self.feed_forward(hidden_states)
713
+ hidden_states = residual + hidden_states
714
+
715
+ outputs = (hidden_states,)
716
+
717
+ if output_attentions:
718
+ outputs += (self_attn_weights,)
719
+
720
+ if use_cache:
721
+ outputs += (present_key_value,)
722
+
723
+ return outputs
724
+
725
+
726
+ InternLM2_START_DOCSTRING = r"""
727
+ This model inherits from [`PreTrainedModel`]. Check the superclass documentation for the generic methods the
728
+ library implements for all its model (such as downloading or saving, resizing the input embeddings, pruning heads
729
+ etc.)
730
+
731
+ This model is also a PyTorch [torch.nn.Module](https://pytorch.org/docs/stable/nn.html#torch.nn.Module) subclass.
732
+ Use it as a regular PyTorch Module and refer to the PyTorch documentation for all matter related to general usage
733
+ and behavior.
734
+
735
+ Parameters:
736
+ config ([`InternLM2Config`]):
737
+ Model configuration class with all the parameters of the model. Initializing with a config file does not
738
+ load the weights associated with the model, only the configuration. Check out the
739
+ [`~PreTrainedModel.from_pretrained`] method to load the model weights.
740
+ """
741
+
742
+
743
+ # Copied from transformers.models.llama.modeling_llama.LlamaPreTrainedModel with Llama->InternLM2
744
+ @add_start_docstrings(
745
+ 'The bare InternLM2 Model outputting raw hidden-states without any specific head on top.',
746
+ InternLM2_START_DOCSTRING,
747
+ )
748
+ class InternLM2PreTrainedModel(PreTrainedModel):
749
+ config_class = InternLM2Config
750
+ base_model_prefix = 'model'
751
+ supports_gradient_checkpointing = True
752
+ _no_split_modules = ['InternLM2DecoderLayer']
753
+ _skip_keys_device_placement = 'past_key_values'
754
+
755
+ def _init_weights(self, module):
756
+ std = self.config.initializer_range
757
+ if isinstance(module, nn.Linear):
758
+ module.weight.data.normal_(mean=0.0, std=std)
759
+ if module.bias is not None:
760
+ module.bias.data.zero_()
761
+ elif isinstance(module, nn.Embedding):
762
+ module.weight.data.normal_(mean=0.0, std=std)
763
+ if module.padding_idx is not None:
764
+ module.weight.data[module.padding_idx].zero_()
765
+
766
+
767
+ InternLM2_INPUTS_DOCSTRING = r"""
768
+ Args:
769
+ input_ids (`torch.LongTensor` of shape `(batch_size, sequence_length)`):
770
+ Indices of input sequence tokens in the vocabulary. Padding will be ignored by default should you provide
771
+ it.
772
+
773
+ Indices can be obtained using [`AutoTokenizer`]. See [`PreTrainedTokenizer.encode`] and
774
+ [`PreTrainedTokenizer.__call__`] for details.
775
+
776
+ [What are input IDs?](../glossary#input-ids)
777
+ attention_mask (`torch.Tensor` of shape `(batch_size, sequence_length)`, *optional*):
778
+ Mask to avoid performing attention on padding token indices. Mask values selected in `[0, 1]`:
779
+
780
+ - 1 for tokens that are **not masked**,
781
+ - 0 for tokens that are **masked**.
782
+
783
+ [What are attention masks?](../glossary#attention-mask)
784
+
785
+ Indices can be obtained using [`AutoTokenizer`]. See [`PreTrainedTokenizer.encode`] and
786
+ [`PreTrainedTokenizer.__call__`] for details.
787
+
788
+ If `past_key_values` is used, optionally only the last `input_ids` have to be input (see
789
+ `past_key_values`).
790
+
791
+ If you want to change padding behavior, you should read [`modeling_opt._prepare_decoder_attention_mask`]
792
+ and modify to your needs. See diagram 1 in [the paper](https://arxiv.org/abs/1910.13461) for more
793
+ information on the default strategy.
794
+
795
+ - 1 indicates the head is **not masked**,
796
+ - 0 indicates the head is **masked**.
797
+ position_ids (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):
798
+ Indices of positions of each input sequence tokens in the position embeddings. Selected in the range `[0,
799
+ config.n_positions - 1]`.
800
+
801
+ [What are position IDs?](../glossary#position-ids)
802
+ past_key_values (`tuple(tuple(torch.FloatTensor))`, *optional*, returned when `use_cache=True` is passed or
803
+ when `config.use_cache=True`):
804
+ Tuple of `tuple(torch.FloatTensor)` of length `config.n_layers`, with each tuple having 2 tensors of shape
805
+ `(batch_size, num_heads, sequence_length, embed_size_per_head)`) and 2 additional tensors of shape
806
+ `(batch_size, num_heads, decoder_sequence_length, embed_size_per_head)`.
807
+
808
+ Contains pre-computed hidden-states (key and values in the self-attention blocks and in the cross-attention
809
+ blocks) that can be used (see `past_key_values` input) to speed up sequential decoding.
810
+
811
+ If `past_key_values` are used, the user can optionally input only the last `input_ids` (those that don't
812
+ have their past key value states given to this model) of shape `(batch_size, 1)` instead of all `input_ids`
813
+ of shape `(batch_size, sequence_length)`.
814
+ inputs_embeds (`torch.FloatTensor` of shape `(batch_size, sequence_length, hidden_size)`, *optional*):
815
+ Optionally, instead of passing `input_ids` you can choose to directly pass an embedded representation. This
816
+ is useful if you want more control over how to convert `input_ids` indices into associated vectors than the
817
+ model's internal embedding lookup matrix.
818
+ use_cache (`bool`, *optional*):
819
+ If set to `True`, `past_key_values` key value states are returned and can be used to speed up decoding (see
820
+ `past_key_values`).
821
+ output_attentions (`bool`, *optional*):
822
+ Whether or not to return the attentions tensors of all attention layers. See `attentions` under returned
823
+ tensors for more detail.
824
+ output_hidden_states (`bool`, *optional*):
825
+ Whether or not to return the hidden states of all layers. See `hidden_states` under returned tensors for
826
+ more detail.
827
+ return_dict (`bool`, *optional*):
828
+ Whether or not to return a [`~utils.ModelOutput`] instead of a plain tuple.
829
+ """
830
+
831
+
832
+ # Modified from transformers.model.llama.modeling_llama.LlamaModel
833
+ @add_start_docstrings(
834
+ 'The bare InternLM2 Model outputting raw hidden-states without any specific head on top.',
835
+ InternLM2_START_DOCSTRING,
836
+ )
837
+ class InternLM2Model(InternLM2PreTrainedModel):
838
+ """
839
+ Transformer decoder consisting of *config.num_hidden_layers* layers. Each layer is a [`InternLM2DecoderLayer`]
840
+
841
+ Args:
842
+ config: InternLM2Config
843
+ """
844
+
845
+ _auto_class = 'AutoModel'
846
+
847
+ def __init__(self, config: InternLM2Config):
848
+ super().__init__(config)
849
+ self.padding_idx = config.pad_token_id
850
+ self.vocab_size = config.vocab_size
851
+ self.config = config
852
+ if not has_flash_attn:
853
+ self.config.attn_implementation = 'eager'
854
+ print('Warning: Flash attention is not available, using eager attention instead.')
855
+
856
+ self.tok_embeddings = nn.Embedding(config.vocab_size, config.hidden_size, self.padding_idx)
857
+
858
+ self.layers = nn.ModuleList([InternLM2DecoderLayer(config) for _ in range(config.num_hidden_layers)])
859
+ self.norm = InternLM2RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
860
+
861
+ self.gradient_checkpointing = False
862
+ # Initialize weights and apply final processing
863
+ self.post_init()
864
+
865
+ def get_input_embeddings(self):
866
+ return self.tok_embeddings
867
+
868
+ def set_input_embeddings(self, value):
869
+ self.tok_embeddings = value
870
+
871
+ def _prepare_decoder_attention_mask(self, attention_mask, input_shape, inputs_embeds, past_key_values_length):
872
+ # create causal mask
873
+ # [bsz, seq_len] -> [bsz, 1, tgt_seq_len, src_seq_len]
874
+ combined_attention_mask = None
875
+ if input_shape[-1] > 1:
876
+ combined_attention_mask = _make_causal_mask(
877
+ input_shape,
878
+ inputs_embeds.dtype,
879
+ device=inputs_embeds.device,
880
+ past_key_values_length=past_key_values_length,
881
+ )
882
+
883
+ if attention_mask is not None:
884
+ # [bsz, seq_len] -> [bsz, 1, tgt_seq_len, src_seq_len]
885
+ expanded_attn_mask = _expand_mask(attention_mask, inputs_embeds.dtype, tgt_len=input_shape[-1]).to(
886
+ inputs_embeds.device
887
+ )
888
+ combined_attention_mask = (
889
+ expanded_attn_mask if combined_attention_mask is None else expanded_attn_mask + combined_attention_mask
890
+ )
891
+
892
+ return combined_attention_mask
893
+
894
+ @add_start_docstrings_to_model_forward(InternLM2_INPUTS_DOCSTRING)
895
+ def forward(
896
+ self,
897
+ input_ids: torch.LongTensor = None,
898
+ attention_mask: Optional[torch.Tensor] = None,
899
+ position_ids: Optional[torch.LongTensor] = None,
900
+ past_key_values: Optional[List[torch.FloatTensor]] = None,
901
+ inputs_embeds: Optional[torch.FloatTensor] = None,
902
+ use_cache: Optional[bool] = None,
903
+ output_attentions: Optional[bool] = None,
904
+ output_hidden_states: Optional[bool] = None,
905
+ return_dict: Optional[bool] = None,
906
+ ) -> Union[Tuple, BaseModelOutputWithPast]:
907
+ output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions
908
+ output_hidden_states = (
909
+ output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
910
+ )
911
+ use_cache = use_cache if use_cache is not None else self.config.use_cache
912
+
913
+ return_dict = return_dict if return_dict is not None else self.config.use_return_dict
914
+
915
+ if self.config.attn_implementation == 'flash_attention_2':
916
+ _import_flash_attn()
917
+
918
+ # retrieve input_ids and inputs_embeds
919
+ if input_ids is not None and inputs_embeds is not None:
920
+ raise ValueError('You cannot specify both input_ids and inputs_embeds at the same time')
921
+ elif input_ids is not None:
922
+ batch_size, seq_length = input_ids.shape[:2]
923
+ elif inputs_embeds is not None:
924
+ batch_size, seq_length = inputs_embeds.shape[:2]
925
+ else:
926
+ raise ValueError('You have to specify either input_ids or inputs_embeds')
927
+
928
+ seq_length_with_past = seq_length
929
+ past_key_values_length = 0
930
+ if past_key_values is not None:
931
+ past_key_values_length = past_key_values[0][0].shape[2]
932
+ seq_length_with_past = seq_length_with_past + past_key_values_length
933
+
934
+ if position_ids is None:
935
+ device = input_ids.device if input_ids is not None else inputs_embeds.device
936
+ position_ids = torch.arange(
937
+ past_key_values_length, seq_length + past_key_values_length, dtype=torch.long, device=device
938
+ )
939
+ position_ids = position_ids.unsqueeze(0)
940
+
941
+ if inputs_embeds is None:
942
+ inputs_embeds = self.tok_embeddings(input_ids)
943
+
944
+ if self.config.attn_implementation == 'flash_attention_2':
945
+ # 2d mask is passed through the layers
946
+ attention_mask = attention_mask if (attention_mask is not None and 0 in attention_mask) else None
947
+ else:
948
+ if attention_mask is None:
949
+ attention_mask = torch.ones(
950
+ (batch_size, seq_length_with_past), dtype=torch.bool, device=inputs_embeds.device
951
+ )
952
+ attention_mask = self._prepare_decoder_attention_mask(
953
+ attention_mask, (batch_size, seq_length), inputs_embeds, past_key_values_length
954
+ )
955
+
956
+ # embed positions
957
+ hidden_states = inputs_embeds
958
+
959
+ if self.gradient_checkpointing and self.training:
960
+ if use_cache:
961
+ logger.warning_once(
962
+ '`use_cache=True` is incompatible with gradient checkpointing. Setting `use_cache=False`...'
963
+ )
964
+ use_cache = False
965
+
966
+ # decoder layers
967
+ all_hidden_states = () if output_hidden_states else None
968
+ all_self_attns = () if output_attentions else None
969
+ next_decoder_cache = () if use_cache else None
970
+
971
+ for idx, decoder_layer in enumerate(self.layers):
972
+ if output_hidden_states:
973
+ all_hidden_states += (hidden_states,)
974
+
975
+ past_key_value = past_key_values[idx] if past_key_values is not None else None
976
+
977
+ if self.gradient_checkpointing and self.training:
978
+
979
+ def create_custom_forward(module):
980
+ def custom_forward(*inputs):
981
+ # None for past_key_value
982
+ return module(*inputs, output_attentions, None)
983
+
984
+ return custom_forward
985
+
986
+ layer_outputs = torch.utils.checkpoint.checkpoint(
987
+ create_custom_forward(decoder_layer),
988
+ hidden_states,
989
+ attention_mask,
990
+ position_ids,
991
+ None,
992
+ )
993
+ else:
994
+ layer_outputs = decoder_layer(
995
+ hidden_states,
996
+ attention_mask=attention_mask,
997
+ position_ids=position_ids,
998
+ past_key_value=past_key_value,
999
+ output_attentions=output_attentions,
1000
+ use_cache=use_cache,
1001
+ )
1002
+
1003
+ hidden_states = layer_outputs[0]
1004
+
1005
+ if use_cache:
1006
+ next_decoder_cache += (layer_outputs[2 if output_attentions else 1],)
1007
+
1008
+ if output_attentions:
1009
+ all_self_attns += (layer_outputs[1],)
1010
+
1011
+ hidden_states = self.norm(hidden_states)
1012
+
1013
+ # add hidden states from the last decoder layer
1014
+ if output_hidden_states:
1015
+ all_hidden_states += (hidden_states,)
1016
+
1017
+ next_cache = next_decoder_cache if use_cache else None
1018
+ if not return_dict:
1019
+ return tuple(v for v in [hidden_states, next_cache, all_hidden_states, all_self_attns] if v is not None)
1020
+ return BaseModelOutputWithPast(
1021
+ last_hidden_state=hidden_states,
1022
+ past_key_values=next_cache,
1023
+ hidden_states=all_hidden_states,
1024
+ attentions=all_self_attns,
1025
+ )
1026
+
1027
+
1028
+ # Modified from transformers.model.llama.modeling_llama.LlamaForCausalLM
1029
+ class InternLM2ForCausalLM(InternLM2PreTrainedModel):
1030
+ _auto_class = 'AutoModelForCausalLM'
1031
+
1032
+ _tied_weights_keys = ['output.weight']
1033
+
1034
+ def __init__(self, config):
1035
+ super().__init__(config)
1036
+ self.model = InternLM2Model(config)
1037
+ self.vocab_size = config.vocab_size
1038
+ self.output = nn.Linear(config.hidden_size, config.vocab_size, bias=False)
1039
+
1040
+ # Initialize weights and apply final processing
1041
+ self.post_init()
1042
+
1043
+ def get_input_embeddings(self):
1044
+ return self.model.tok_embeddings
1045
+
1046
+ def set_input_embeddings(self, value):
1047
+ self.model.tok_embeddings = value
1048
+
1049
+ def get_output_embeddings(self):
1050
+ return self.output
1051
+
1052
+ def set_output_embeddings(self, new_embeddings):
1053
+ self.output = new_embeddings
1054
+
1055
+ def set_decoder(self, decoder):
1056
+ self.model = decoder
1057
+
1058
+ def get_decoder(self):
1059
+ return self.model
1060
+
1061
+ @add_start_docstrings_to_model_forward(InternLM2_INPUTS_DOCSTRING)
1062
+ @replace_return_docstrings(output_type=CausalLMOutputWithPast, config_class=_CONFIG_FOR_DOC)
1063
+ def forward(
1064
+ self,
1065
+ input_ids: torch.LongTensor = None,
1066
+ attention_mask: Optional[torch.Tensor] = None,
1067
+ position_ids: Optional[torch.LongTensor] = None,
1068
+ past_key_values: Optional[List[torch.FloatTensor]] = None,
1069
+ inputs_embeds: Optional[torch.FloatTensor] = None,
1070
+ labels: Optional[torch.LongTensor] = None,
1071
+ use_cache: Optional[bool] = None,
1072
+ output_attentions: Optional[bool] = None,
1073
+ output_hidden_states: Optional[bool] = None,
1074
+ return_dict: Optional[bool] = None,
1075
+ ) -> Union[Tuple, CausalLMOutputWithPast]:
1076
+ r"""
1077
+ Args:
1078
+ labels (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):
1079
+ Labels for computing the masked language modeling loss. Indices should either be in `[0, ...,
1080
+ config.vocab_size]` or -100 (see `input_ids` docstring). Tokens with indices set to `-100` are ignored
1081
+ (masked), the loss is only computed for the tokens with labels in `[0, ..., config.vocab_size]`.
1082
+
1083
+ Returns:
1084
+
1085
+ Example:
1086
+
1087
+ ```python
1088
+ >>> from transformers import AutoTokenizer, InternLM2ForCausalLM
1089
+
1090
+ >>> model = InternLM2ForCausalLM.from_pretrained(PATH_TO_CONVERTED_WEIGHTS)
1091
+ >>> tokenizer = AutoTokenizer.from_pretrained(PATH_TO_CONVERTED_TOKENIZER)
1092
+
1093
+ >>> prompt = "Hey, are you conscious? Can you talk to me?"
1094
+ >>> inputs = tokenizer(prompt, return_tensors="pt")
1095
+
1096
+ >>> # Generate
1097
+ >>> generate_ids = model.generate(inputs.input_ids, max_length=30)
1098
+ >>> tokenizer.batch_decode(generate_ids, skip_special_tokens=True, clean_up_tokenization_spaces=False)[0]
1099
+ "Hey, are you conscious? Can you talk to me?\nI'm not conscious, but I can talk to you."
1100
+ ```"""
1101
+
1102
+ output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions
1103
+ output_hidden_states = (
1104
+ output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
1105
+ )
1106
+ return_dict = return_dict if return_dict is not None else self.config.use_return_dict
1107
+
1108
+ # decoder outputs consists of (dec_features, layer_state, dec_hidden, dec_attn)
1109
+ outputs = self.model(
1110
+ input_ids=input_ids,
1111
+ attention_mask=attention_mask,
1112
+ position_ids=position_ids,
1113
+ past_key_values=past_key_values,
1114
+ inputs_embeds=inputs_embeds,
1115
+ use_cache=use_cache,
1116
+ output_attentions=output_attentions,
1117
+ output_hidden_states=output_hidden_states,
1118
+ return_dict=return_dict,
1119
+ )
1120
+
1121
+ hidden_states = outputs[0]
1122
+ logits = self.output(hidden_states)
1123
+ logits = logits.float()
1124
+
1125
+ loss = None
1126
+ if labels is not None:
1127
+ # Shift so that tokens < n predict n
1128
+ shift_logits = logits[..., :-1, :].contiguous()
1129
+ shift_labels = labels[..., 1:].contiguous()
1130
+ # Flatten the tokens
1131
+ loss_fct = CrossEntropyLoss()
1132
+ shift_logits = shift_logits.view(-1, self.config.vocab_size)
1133
+ shift_labels = shift_labels.view(-1)
1134
+ # Enable model parallelism
1135
+ shift_labels = shift_labels.to(shift_logits.device)
1136
+ loss = loss_fct(shift_logits, shift_labels)
1137
+
1138
+ if not return_dict:
1139
+ output = (logits,) + outputs[1:]
1140
+ return (loss,) + output if loss is not None else output
1141
+
1142
+ device = input_ids.device if input_ids is not None else inputs_embeds.device
1143
+ output = CausalLMOutputWithPast(
1144
+ loss=loss,
1145
+ logits=logits,
1146
+ past_key_values=outputs.past_key_values,
1147
+ hidden_states=outputs.hidden_states,
1148
+ attentions=outputs.attentions,
1149
+ )
1150
+ output['logits'] = output['logits'].to(device)
1151
+ return output
1152
+
1153
+ def prepare_inputs_for_generation(
1154
+ self, input_ids, past_key_values=None, attention_mask=None, inputs_embeds=None, **kwargs
1155
+ ):
1156
+ if past_key_values is not None:
1157
+ past_length = past_key_values[0][0].shape[2]
1158
+
1159
+ # Some generation methods already pass only the last input ID
1160
+ if input_ids.shape[1] > past_length:
1161
+ remove_prefix_length = past_length
1162
+ else:
1163
+ # Default to old behavior: keep only final ID
1164
+ remove_prefix_length = input_ids.shape[1] - 1
1165
+
1166
+ input_ids = input_ids[:, remove_prefix_length:]
1167
+
1168
+ position_ids = kwargs.get('position_ids', None)
1169
+ if attention_mask is not None and position_ids is None:
1170
+ # create position_ids on the fly for batch generation
1171
+ position_ids = attention_mask.long().cumsum(-1) - 1
1172
+ position_ids.masked_fill_(attention_mask == 0, 1)
1173
+ if past_key_values:
1174
+ position_ids = position_ids[:, -input_ids.shape[1]:]
1175
+
1176
+ # if `inputs_embeds` are passed, we only want to use them in the 1st generation step
1177
+ if inputs_embeds is not None and past_key_values is None:
1178
+ model_inputs = {'inputs_embeds': inputs_embeds}
1179
+ else:
1180
+ model_inputs = {'input_ids': input_ids}
1181
+
1182
+ model_inputs.update(
1183
+ {
1184
+ 'position_ids': position_ids,
1185
+ 'past_key_values': past_key_values,
1186
+ 'use_cache': kwargs.get('use_cache'),
1187
+ 'attention_mask': attention_mask,
1188
+ }
1189
+ )
1190
+ return model_inputs
1191
+
1192
+ @staticmethod
1193
+ def _reorder_cache(past_key_values, beam_idx):
1194
+ reordered_past = ()
1195
+ for layer_past in past_key_values:
1196
+ reordered_past += (
1197
+ tuple(past_state.index_select(0, beam_idx.to(past_state.device)) for past_state in layer_past),
1198
+ )
1199
+ return reordered_past
1200
+
1201
+ def build_inputs(self, tokenizer, query: str, history: List[Tuple[str, str]] = [], meta_instruction=''):
1202
+ if tokenizer.add_bos_token:
1203
+ prompt = ''
1204
+ else:
1205
+ prompt = tokenizer.bos_token
1206
+ if meta_instruction:
1207
+ prompt += f"""<|im_start|>system\n{meta_instruction}<|im_end|>\n"""
1208
+ for record in history:
1209
+ prompt += f"""<|im_start|>user\n{record[0]}<|im_end|>\n<|im_start|>assistant\n{record[1]}<|im_end|>\n"""
1210
+ prompt += f"""<|im_start|>user\n{query}<|im_end|>\n<|im_start|>assistant\n"""
1211
+ return tokenizer([prompt], return_tensors='pt')
1212
+
1213
+ @torch.no_grad()
1214
+ def chat(
1215
+ self,
1216
+ tokenizer,
1217
+ query: str,
1218
+ history: List[Tuple[str, str]] = [],
1219
+ streamer: Optional[BaseStreamer] = None,
1220
+ max_new_tokens: int = 1024,
1221
+ do_sample: bool = True,
1222
+ temperature: float = 0.8,
1223
+ top_p: float = 0.8,
1224
+ meta_instruction: str = 'You are an AI assistant whose name is InternLM (书生·浦语).\n'
1225
+ '- InternLM (书生·浦语) is a conversational language model that is developed by Shanghai AI Laboratory (上海人工智能实验室). It is designed to be helpful, honest, and harmless.\n'
1226
+ '- InternLM (书生·浦语) can understand and communicate fluently in the language chosen by the user such as English and 中文.',
1227
+ **kwargs,
1228
+ ):
1229
+ inputs = self.build_inputs(tokenizer, query, history, meta_instruction)
1230
+ inputs = {k: v.to(self.device) for k, v in inputs.items() if torch.is_tensor(v)}
1231
+ # also add end-of-assistant token in eos token id to avoid unnecessary generation
1232
+ eos_token_id = [tokenizer.eos_token_id, tokenizer.convert_tokens_to_ids(['<|im_end|>'])[0]]
1233
+ outputs = self.generate(
1234
+ **inputs,
1235
+ streamer=streamer,
1236
+ max_new_tokens=max_new_tokens,
1237
+ do_sample=do_sample,
1238
+ temperature=temperature,
1239
+ top_p=top_p,
1240
+ eos_token_id=eos_token_id,
1241
+ **kwargs,
1242
+ )
1243
+ outputs = outputs[0].cpu().tolist()[len(inputs['input_ids'][0]):]
1244
+ response = tokenizer.decode(outputs, skip_special_tokens=True)
1245
+ response = response.split('<|im_end|>')[0]
1246
+ history = history + [(query, response)]
1247
+ return response, history
1248
+
1249
+ @torch.no_grad()
1250
+ def stream_chat(
1251
+ self,
1252
+ tokenizer,
1253
+ query: str,
1254
+ history: List[Tuple[str, str]] = [],
1255
+ max_new_tokens: int = 1024,
1256
+ do_sample: bool = True,
1257
+ temperature: float = 0.8,
1258
+ top_p: float = 0.8,
1259
+ **kwargs,
1260
+ ):
1261
+ """
1262
+ Return a generator in format: (response, history)
1263
+ Eg.
1264
+ ('你好,有什么可以帮助您的吗', [('你好', '你好,有什么可以帮助您的吗')])
1265
+ ('你好,有什么可以帮助您的吗?', [('你好', '你好,有什么可以帮助您的吗?')])
1266
+ """
1267
+ if BaseStreamer is None:
1268
+ raise ModuleNotFoundError(
1269
+ 'The version of `transformers` is too low. Please make sure '
1270
+ 'that you have installed `transformers>=4.28.0`.'
1271
+ )
1272
+
1273
+ response_queue = queue.Queue(maxsize=20)
1274
+
1275
+ class ChatStreamer(BaseStreamer):
1276
+ def __init__(self, tokenizer) -> None:
1277
+ super().__init__()
1278
+ self.tokenizer = tokenizer
1279
+ self.queue = response_queue
1280
+ self.query = query
1281
+ self.history = history
1282
+ self.response = ''
1283
+ self.cache = []
1284
+ self.received_inputs = False
1285
+ self.queue.put((self.response, history + [(self.query, self.response)]))
1286
+
1287
+ def put(self, value):
1288
+ if len(value.shape) > 1 and value.shape[0] > 1:
1289
+ raise ValueError('ChatStreamer only supports batch size 1')
1290
+ elif len(value.shape) > 1:
1291
+ value = value[0]
1292
+
1293
+ if not self.received_inputs:
1294
+ # The first received value is input_ids, ignore here
1295
+ self.received_inputs = True
1296
+ return
1297
+
1298
+ self.cache.extend(value.tolist())
1299
+ token = self.tokenizer.decode(self.cache, skip_special_tokens=True)
1300
+ if token.strip() != '<|im_end|>':
1301
+ self.response = self.response + token
1302
+ history = self.history + [(self.query, self.response)]
1303
+ self.queue.put((self.response, history))
1304
+ self.cache = []
1305
+ else:
1306
+ self.end()
1307
+
1308
+ def end(self):
1309
+ self.queue.put(None)
1310
+
1311
+ def stream_producer():
1312
+ return self.chat(
1313
+ tokenizer=tokenizer,
1314
+ query=query,
1315
+ streamer=ChatStreamer(tokenizer=tokenizer),
1316
+ history=history,
1317
+ max_new_tokens=max_new_tokens,
1318
+ do_sample=do_sample,
1319
+ temperature=temperature,
1320
+ top_p=top_p,
1321
+ **kwargs,
1322
+ )
1323
+
1324
+ def consumer():
1325
+ producer = threading.Thread(target=stream_producer)
1326
+ producer.start()
1327
+ while True:
1328
+ res = response_queue.get()
1329
+ if res is None:
1330
+ return
1331
+ yield res
1332
+
1333
+ return consumer()
1334
+
1335
+
1336
+ # Copied from transformers.model.llama.modeling_llama.LlamaForSequenceClassification with Llama->InternLM2
1337
+ @add_start_docstrings(
1338
+ """
1339
+ The InternLM2 Model transformer with a sequence classification head on top (linear layer).
1340
+
1341
+ [`InternLM2ForSequenceClassification`] uses the last token in order to do the classification,
1342
+ as other causal models (e.g. GPT-2) do.
1343
+
1344
+ Since it does classification on the last token, it requires to know the position of the last token. If a
1345
+ `pad_token_id` is defined in the configuration, it finds the last token that is not a padding token in each row. If
1346
+ no `pad_token_id` is defined, it simply takes the last value in each row of the batch. Since it cannot guess the
1347
+ padding tokens when `inputs_embeds` are passed instead of `input_ids`, it does the same (take the last value in
1348
+ each row of the batch).
1349
+ """,
1350
+ InternLM2_START_DOCSTRING,
1351
+ )
1352
+ class InternLM2ForSequenceClassification(InternLM2PreTrainedModel):
1353
+ def __init__(self, config):
1354
+ super().__init__(config)
1355
+ self.num_labels = config.num_labels
1356
+ self.model = InternLM2Model(config)
1357
+ self.score = nn.Linear(config.hidden_size, self.num_labels, bias=False)
1358
+
1359
+ # Initialize weights and apply final processing
1360
+ self.post_init()
1361
+
1362
+ def get_input_embeddings(self):
1363
+ return self.model.tok_embeddings
1364
+
1365
+ def set_input_embeddings(self, value):
1366
+ self.model.tok_embeddings = value
1367
+
1368
+ @add_start_docstrings_to_model_forward(InternLM2_INPUTS_DOCSTRING)
1369
+ def forward(
1370
+ self,
1371
+ input_ids: torch.LongTensor = None,
1372
+ attention_mask: Optional[torch.Tensor] = None,
1373
+ position_ids: Optional[torch.LongTensor] = None,
1374
+ past_key_values: Optional[List[torch.FloatTensor]] = None,
1375
+ inputs_embeds: Optional[torch.FloatTensor] = None,
1376
+ labels: Optional[torch.LongTensor] = None,
1377
+ use_cache: Optional[bool] = None,
1378
+ output_attentions: Optional[bool] = None,
1379
+ output_hidden_states: Optional[bool] = None,
1380
+ return_dict: Optional[bool] = None,
1381
+ ) -> Union[Tuple, SequenceClassifierOutputWithPast]:
1382
+ r"""
1383
+ labels (`torch.LongTensor` of shape `(batch_size,)`, *optional*):
1384
+ Labels for computing the sequence classification/regression loss. Indices should be in `[0, ...,
1385
+ config.num_labels - 1]`. If `config.num_labels == 1` a regression loss is computed (Mean-Square loss), If
1386
+ `config.num_labels > 1` a classification loss is computed (Cross-Entropy).
1387
+ """
1388
+ return_dict = return_dict if return_dict is not None else self.config.use_return_dict
1389
+
1390
+ transformer_outputs = self.model(
1391
+ input_ids,
1392
+ attention_mask=attention_mask,
1393
+ position_ids=position_ids,
1394
+ past_key_values=past_key_values,
1395
+ inputs_embeds=inputs_embeds,
1396
+ use_cache=use_cache,
1397
+ output_attentions=output_attentions,
1398
+ output_hidden_states=output_hidden_states,
1399
+ return_dict=return_dict,
1400
+ )
1401
+ hidden_states = transformer_outputs[0]
1402
+ logits = self.score(hidden_states)
1403
+
1404
+ if input_ids is not None:
1405
+ batch_size = input_ids.shape[0]
1406
+ else:
1407
+ batch_size = inputs_embeds.shape[0]
1408
+
1409
+ if self.config.pad_token_id is None and batch_size != 1:
1410
+ raise ValueError('Cannot handle batch sizes > 1 if no padding token is defined.')
1411
+ if self.config.pad_token_id is None:
1412
+ sequence_lengths = -1
1413
+ else:
1414
+ if input_ids is not None:
1415
+ sequence_lengths = (torch.eq(input_ids, self.config.pad_token_id).int().argmax(-1) - 1).to(
1416
+ logits.device
1417
+ )
1418
+ else:
1419
+ sequence_lengths = -1
1420
+
1421
+ pooled_logits = logits[torch.arange(batch_size, device=logits.device), sequence_lengths]
1422
+
1423
+ loss = None
1424
+ if labels is not None:
1425
+ labels = labels.to(logits.device)
1426
+ if self.config.problem_type is None:
1427
+ if self.num_labels == 1:
1428
+ self.config.problem_type = 'regression'
1429
+ elif self.num_labels > 1 and (labels.dtype == torch.long or labels.dtype == torch.int):
1430
+ self.config.problem_type = 'single_label_classification'
1431
+ else:
1432
+ self.config.problem_type = 'multi_label_classification'
1433
+
1434
+ if self.config.problem_type == 'regression':
1435
+ loss_fct = MSELoss()
1436
+ if self.num_labels == 1:
1437
+ loss = loss_fct(pooled_logits.squeeze(), labels.squeeze())
1438
+ else:
1439
+ loss = loss_fct(pooled_logits, labels)
1440
+ elif self.config.problem_type == 'single_label_classification':
1441
+ loss_fct = CrossEntropyLoss()
1442
+ loss = loss_fct(pooled_logits.view(-1, self.num_labels), labels.view(-1))
1443
+ elif self.config.problem_type == 'multi_label_classification':
1444
+ loss_fct = BCEWithLogitsLoss()
1445
+ loss = loss_fct(pooled_logits, labels)
1446
+ if not return_dict:
1447
+ output = (pooled_logits,) + transformer_outputs[1:]
1448
+ return ((loss,) + output) if loss is not None else output
1449
+
1450
+ return SequenceClassifierOutputWithPast(
1451
+ loss=loss,
1452
+ logits=pooled_logits,
1453
+ past_key_values=transformer_outputs.past_key_values,
1454
+ hidden_states=transformer_outputs.hidden_states,
1455
+ attentions=transformer_outputs.attentions,
1456
+ )
native_backbone/modeling_internvl_chat.py ADDED
@@ -0,0 +1,363 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # --------------------------------------------------------
2
+ # InternVL
3
+ # Copyright (c) 2024 OpenGVLab
4
+ # Licensed under The MIT License [see LICENSE for details]
5
+ # --------------------------------------------------------
6
+ import warnings
7
+ from typing import Any, List, Optional, Tuple, Union
8
+
9
+ import torch.utils.checkpoint
10
+ import transformers
11
+ from torch import nn
12
+ from torch.nn import CrossEntropyLoss
13
+ from transformers import (AutoModel, GenerationConfig, LlamaForCausalLM,
14
+ LlamaTokenizer, Qwen2ForCausalLM)
15
+ from transformers.modeling_outputs import CausalLMOutputWithPast
16
+ from transformers.modeling_utils import PreTrainedModel
17
+ from transformers.utils import ModelOutput, logging
18
+
19
+ from .configuration_internvl_chat import InternVLChatConfig
20
+ from .conversation import get_conv_template
21
+ from .modeling_intern_vit import InternVisionModel, has_flash_attn
22
+ from .modeling_internlm2 import InternLM2ForCausalLM
23
+
24
+ logger = logging.get_logger(__name__)
25
+
26
+
27
+ def version_cmp(v1, v2, op='eq'):
28
+ import operator
29
+
30
+ from packaging import version
31
+ op_func = getattr(operator, op)
32
+ return op_func(version.parse(v1), version.parse(v2))
33
+
34
+
35
+ class InternVLChatModel(PreTrainedModel):
36
+ config_class = InternVLChatConfig
37
+ main_input_name = 'pixel_values'
38
+ _supports_flash_attn_2 = True
39
+ supports_gradient_checkpointing = True
40
+ _no_split_modules = ['InternVisionModel', 'LlamaDecoderLayer', 'InternLM2DecoderLayer',
41
+ 'Qwen2DecoderLayer']
42
+
43
+ def __init__(self, config: InternVLChatConfig, vision_model=None, language_model=None, use_flash_attn=True):
44
+ super().__init__(config)
45
+
46
+ assert version_cmp(transformers.__version__, '4.36.2', 'ge')
47
+ image_size = config.force_image_size or config.vision_config.image_size
48
+ patch_size = config.vision_config.patch_size
49
+ self.patch_size = patch_size
50
+ self.select_layer = config.select_layer
51
+ self.template = config.template
52
+ self.num_image_token = int((image_size // patch_size) ** 2 * (config.downsample_ratio ** 2))
53
+ self.downsample_ratio = config.downsample_ratio
54
+ self.ps_version = config.ps_version
55
+ use_flash_attn = use_flash_attn if has_flash_attn else False
56
+ config.vision_config.use_flash_attn = True if use_flash_attn else False
57
+ config.llm_config.attn_implementation = 'flash_attention_2' if use_flash_attn else 'eager'
58
+
59
+ logger.info(f'num_image_token: {self.num_image_token}')
60
+ logger.info(f'ps_version: {self.ps_version}')
61
+ if vision_model is not None:
62
+ self.vision_model = vision_model
63
+ else:
64
+ self.vision_model = InternVisionModel(config.vision_config)
65
+ if language_model is not None:
66
+ self.language_model = language_model
67
+ else:
68
+ if config.llm_config.architectures[0] == 'LlamaForCausalLM':
69
+ self.language_model = LlamaForCausalLM(config.llm_config)
70
+ elif config.llm_config.architectures[0] == 'InternLM2ForCausalLM':
71
+ self.language_model = InternLM2ForCausalLM(config.llm_config)
72
+ elif config.llm_config.architectures[0] == 'Qwen2ForCausalLM':
73
+ self.language_model = Qwen2ForCausalLM(config.llm_config)
74
+ else:
75
+ raise NotImplementedError(f'{config.llm_config.architectures[0]} is not implemented.')
76
+
77
+ vit_hidden_size = config.vision_config.hidden_size
78
+ llm_hidden_size = config.llm_config.hidden_size
79
+
80
+ self.mlp1 = nn.Sequential(
81
+ nn.LayerNorm(vit_hidden_size * int(1 / self.downsample_ratio) ** 2),
82
+ nn.Linear(vit_hidden_size * int(1 / self.downsample_ratio) ** 2, llm_hidden_size),
83
+ nn.GELU(),
84
+ nn.Linear(llm_hidden_size, llm_hidden_size)
85
+ )
86
+
87
+ self.img_context_token_id = None
88
+ self.conv_template = get_conv_template(self.template)
89
+ self.system_message = self.conv_template.system_message
90
+
91
+ def forward(
92
+ self,
93
+ pixel_values: torch.FloatTensor,
94
+ input_ids: torch.LongTensor = None,
95
+ attention_mask: Optional[torch.Tensor] = None,
96
+ position_ids: Optional[torch.LongTensor] = None,
97
+ image_flags: Optional[torch.LongTensor] = None,
98
+ past_key_values: Optional[List[torch.FloatTensor]] = None,
99
+ labels: Optional[torch.LongTensor] = None,
100
+ use_cache: Optional[bool] = None,
101
+ output_attentions: Optional[bool] = None,
102
+ output_hidden_states: Optional[bool] = None,
103
+ return_dict: Optional[bool] = None,
104
+ ) -> Union[Tuple, CausalLMOutputWithPast]:
105
+ return_dict = return_dict if return_dict is not None else self.config.use_return_dict
106
+
107
+ image_flags = image_flags.squeeze(-1)
108
+ input_embeds = self.language_model.get_input_embeddings()(input_ids)
109
+
110
+ vit_embeds = self.extract_feature(pixel_values)
111
+ vit_embeds = vit_embeds[image_flags == 1]
112
+ vit_batch_size = pixel_values.shape[0]
113
+
114
+ B, N, C = input_embeds.shape
115
+ input_embeds = input_embeds.reshape(B * N, C)
116
+
117
+ if torch.distributed.get_rank() == 0:
118
+ print(f'dynamic ViT batch size: {vit_batch_size}, images per sample: {vit_batch_size / B}, dynamic token length: {N}')
119
+
120
+ input_ids = input_ids.reshape(B * N)
121
+ selected = (input_ids == self.img_context_token_id)
122
+ try:
123
+ input_embeds[selected] = input_embeds[selected] * 0.0 + vit_embeds.reshape(-1, C)
124
+ except Exception as e:
125
+ vit_embeds = vit_embeds.reshape(-1, C)
126
+ print(f'warning: {e}, input_embeds[selected].shape={input_embeds[selected].shape}, '
127
+ f'vit_embeds.shape={vit_embeds.shape}')
128
+ n_token = min(selected.sum(), vit_embeds.size(0))
129
+ input_embeds[selected][:n_token] = input_embeds[selected][:n_token] * 0.0 + vit_embeds[:n_token]
130
+
131
+ input_embeds = input_embeds.reshape(B, N, C)
132
+
133
+ outputs = self.language_model(
134
+ inputs_embeds=input_embeds,
135
+ attention_mask=attention_mask,
136
+ position_ids=position_ids,
137
+ past_key_values=past_key_values,
138
+ use_cache=use_cache,
139
+ output_attentions=output_attentions,
140
+ output_hidden_states=output_hidden_states,
141
+ return_dict=return_dict,
142
+ )
143
+ logits = outputs.logits
144
+
145
+ loss = None
146
+ if labels is not None:
147
+ # Shift so that tokens < n predict n
148
+ shift_logits = logits[..., :-1, :].contiguous()
149
+ shift_labels = labels[..., 1:].contiguous()
150
+ # Flatten the tokens
151
+ loss_fct = CrossEntropyLoss()
152
+ shift_logits = shift_logits.view(-1, self.language_model.config.vocab_size)
153
+ shift_labels = shift_labels.view(-1)
154
+ # Enable model parallelism
155
+ shift_labels = shift_labels.to(shift_logits.device)
156
+ loss = loss_fct(shift_logits, shift_labels)
157
+
158
+ if not return_dict:
159
+ output = (logits,) + outputs[1:]
160
+ return (loss,) + output if loss is not None else output
161
+
162
+ return CausalLMOutputWithPast(
163
+ loss=loss,
164
+ logits=logits,
165
+ past_key_values=outputs.past_key_values,
166
+ hidden_states=outputs.hidden_states,
167
+ attentions=outputs.attentions,
168
+ )
169
+
170
+ def pixel_shuffle(self, x, scale_factor=0.5):
171
+ n, w, h, c = x.size()
172
+ # N, W, H, C --> N, W, H * scale, C // scale
173
+ x = x.view(n, w, int(h * scale_factor), int(c / scale_factor))
174
+ # N, W, H * scale, C // scale --> N, H * scale, W, C // scale
175
+ x = x.permute(0, 2, 1, 3).contiguous()
176
+ # N, H * scale, W, C // scale --> N, H * scale, W * scale, C // (scale ** 2)
177
+ x = x.view(n, int(h * scale_factor), int(w * scale_factor),
178
+ int(c / (scale_factor * scale_factor)))
179
+ if self.ps_version == 'v1':
180
+ warnings.warn("In ps_version 'v1', the height and width have not been swapped back, "
181
+ 'which results in a transposed image.')
182
+ else:
183
+ x = x.permute(0, 2, 1, 3).contiguous()
184
+ return x
185
+
186
+ def extract_feature(self, pixel_values):
187
+ if self.select_layer == -1:
188
+ vit_embeds = self.vision_model(
189
+ pixel_values=pixel_values,
190
+ output_hidden_states=False,
191
+ return_dict=True).last_hidden_state
192
+ else:
193
+ vit_embeds = self.vision_model(
194
+ pixel_values=pixel_values,
195
+ output_hidden_states=True,
196
+ return_dict=True).hidden_states[self.select_layer]
197
+ vit_embeds = vit_embeds[:, 1:, :]
198
+
199
+ h = w = int(vit_embeds.shape[1] ** 0.5)
200
+ vit_embeds = vit_embeds.reshape(vit_embeds.shape[0], h, w, -1)
201
+ vit_embeds = self.pixel_shuffle(vit_embeds, scale_factor=self.downsample_ratio)
202
+ vit_embeds = vit_embeds.reshape(vit_embeds.shape[0], -1, vit_embeds.shape[-1])
203
+ vit_embeds = self.mlp1(vit_embeds)
204
+ return vit_embeds
205
+
206
+ def batch_chat(self, tokenizer, pixel_values, questions, generation_config, num_patches_list=None,
207
+ history=None, return_history=False, IMG_START_TOKEN='<img>', IMG_END_TOKEN='</img>',
208
+ IMG_CONTEXT_TOKEN='<IMG_CONTEXT>', verbose=False, image_counts=None):
209
+ if history is not None or return_history:
210
+ print('Now multi-turn chat is not supported in batch_chat.')
211
+ raise NotImplementedError
212
+
213
+ if image_counts is not None:
214
+ num_patches_list = image_counts
215
+ print('Warning: `image_counts` is deprecated. Please use `num_patches_list` instead.')
216
+
217
+ img_context_token_id = tokenizer.convert_tokens_to_ids(IMG_CONTEXT_TOKEN)
218
+ self.img_context_token_id = img_context_token_id
219
+
220
+ if verbose and pixel_values is not None:
221
+ image_bs = pixel_values.shape[0]
222
+ print(f'dynamic ViT batch size: {image_bs}')
223
+
224
+ queries = []
225
+ for idx, num_patches in enumerate(num_patches_list):
226
+ question = questions[idx]
227
+ if pixel_values is not None and '<image>' not in question:
228
+ question = '<image>\n' + question
229
+ template = get_conv_template(self.template)
230
+ template.system_message = self.system_message
231
+ template.append_message(template.roles[0], question)
232
+ template.append_message(template.roles[1], None)
233
+ query = template.get_prompt()
234
+
235
+ image_tokens = IMG_START_TOKEN + IMG_CONTEXT_TOKEN * self.num_image_token * num_patches + IMG_END_TOKEN
236
+ query = query.replace('<image>', image_tokens, 1)
237
+ queries.append(query)
238
+
239
+ tokenizer.padding_side = 'left'
240
+ model_inputs = tokenizer(queries, return_tensors='pt', padding=True)
241
+ input_ids = model_inputs['input_ids'].cuda()
242
+ attention_mask = model_inputs['attention_mask'].cuda()
243
+ eos_token_id = tokenizer.convert_tokens_to_ids(template.sep.strip())
244
+ generation_config['eos_token_id'] = eos_token_id
245
+ generation_output = self.generate(
246
+ pixel_values=pixel_values,
247
+ input_ids=input_ids,
248
+ attention_mask=attention_mask,
249
+ **generation_config
250
+ )
251
+ responses = tokenizer.batch_decode(generation_output, skip_special_tokens=True)
252
+ responses = [response.split(template.sep)[0].strip() for response in responses]
253
+ return responses
254
+
255
+ def chat(self, tokenizer, pixel_values, question, generation_config, history=None, return_history=False,
256
+ num_patches_list=None, IMG_START_TOKEN='<img>', IMG_END_TOKEN='</img>', IMG_CONTEXT_TOKEN='<IMG_CONTEXT>',
257
+ verbose=False):
258
+
259
+ if history is None and pixel_values is not None and '<image>' not in question:
260
+ question = '<image>\n' + question
261
+
262
+ if num_patches_list is None:
263
+ num_patches_list = [pixel_values.shape[0]] if pixel_values is not None else []
264
+ assert pixel_values is None or len(pixel_values) == sum(num_patches_list)
265
+
266
+ img_context_token_id = tokenizer.convert_tokens_to_ids(IMG_CONTEXT_TOKEN)
267
+ self.img_context_token_id = img_context_token_id
268
+
269
+ template = get_conv_template(self.template)
270
+ template.system_message = self.system_message
271
+ eos_token_id = tokenizer.convert_tokens_to_ids(template.sep.strip())
272
+
273
+ history = [] if history is None else history
274
+ for (old_question, old_answer) in history:
275
+ template.append_message(template.roles[0], old_question)
276
+ template.append_message(template.roles[1], old_answer)
277
+ template.append_message(template.roles[0], question)
278
+ template.append_message(template.roles[1], None)
279
+ query = template.get_prompt()
280
+
281
+ if verbose and pixel_values is not None:
282
+ image_bs = pixel_values.shape[0]
283
+ print(f'dynamic ViT batch size: {image_bs}')
284
+
285
+ for num_patches in num_patches_list:
286
+ image_tokens = IMG_START_TOKEN + IMG_CONTEXT_TOKEN * self.num_image_token * num_patches + IMG_END_TOKEN
287
+ query = query.replace('<image>', image_tokens, 1)
288
+
289
+ model_inputs = tokenizer(query, return_tensors='pt')
290
+ input_ids = model_inputs['input_ids'].cuda()
291
+ attention_mask = model_inputs['attention_mask'].cuda()
292
+ generation_config['eos_token_id'] = eos_token_id
293
+ generation_output = self.generate(
294
+ pixel_values=pixel_values,
295
+ input_ids=input_ids,
296
+ attention_mask=attention_mask,
297
+ **generation_config
298
+ )
299
+ response = tokenizer.batch_decode(generation_output, skip_special_tokens=True)[0]
300
+ response = response.split(template.sep.strip())[0].strip()
301
+ history.append((question, response))
302
+ if return_history:
303
+ return response, history
304
+ else:
305
+ query_to_print = query.replace(IMG_CONTEXT_TOKEN, '')
306
+ query_to_print = query_to_print.replace(f'{IMG_START_TOKEN}{IMG_END_TOKEN}', '<image>')
307
+ if verbose:
308
+ print(query_to_print, response)
309
+ return response
310
+
311
+ @torch.no_grad()
312
+ def generate(
313
+ self,
314
+ pixel_values: Optional[torch.FloatTensor] = None,
315
+ input_ids: Optional[torch.FloatTensor] = None,
316
+ attention_mask: Optional[torch.LongTensor] = None,
317
+ visual_features: Optional[torch.FloatTensor] = None,
318
+ generation_config: Optional[GenerationConfig] = None,
319
+ output_hidden_states: Optional[bool] = None,
320
+ return_dict: Optional[bool] = None,
321
+ **generate_kwargs,
322
+ ) -> torch.LongTensor:
323
+
324
+ assert self.img_context_token_id is not None
325
+ if pixel_values is not None:
326
+ if visual_features is not None:
327
+ vit_embeds = visual_features
328
+ else:
329
+ vit_embeds = self.extract_feature(pixel_values)
330
+ input_embeds = self.language_model.get_input_embeddings()(input_ids)
331
+ B, N, C = input_embeds.shape
332
+ input_embeds = input_embeds.reshape(B * N, C)
333
+
334
+ input_ids = input_ids.reshape(B * N)
335
+ selected = (input_ids == self.img_context_token_id)
336
+ assert selected.sum() != 0
337
+ input_embeds[selected] = vit_embeds.reshape(-1, C).to(input_embeds.device)
338
+
339
+ input_embeds = input_embeds.reshape(B, N, C)
340
+ else:
341
+ input_embeds = self.language_model.get_input_embeddings()(input_ids)
342
+
343
+ outputs = self.language_model.generate(
344
+ inputs_embeds=input_embeds,
345
+ attention_mask=attention_mask,
346
+ generation_config=generation_config,
347
+ output_hidden_states=output_hidden_states,
348
+ # return_dict=return_dict, # return_dict is not supported in transformers 4.44
349
+ use_cache=True,
350
+ **generate_kwargs,
351
+ )
352
+
353
+ return outputs
354
+
355
+ @property
356
+ def lm_head(self):
357
+ return self.language_model.get_output_embeddings()
358
+
359
+ def get_input_embeddings(self):
360
+ return self.language_model.get_input_embeddings()
361
+
362
+ def get_output_embeddings(self):
363
+ return self.language_model.get_output_embeddings()
native_backbone/native-00001.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:a8595a8061a1eacd66c8e8872754ff18b36113c660d01ebfdf91cf9701477543
3
+ size 4878234984
native_backbone/native-00002.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:66ebb3671f03e25fc350750513958e39078d851ac332aedb92cbaa5aba071bc6
3
+ size 4882437624
native_backbone/native-00003.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:186bf3a374834e1e573e9d8ff373e1dec200eca2b5dc23e2ddcdcc75cf4ecb5d
3
+ size 4844680424
native_backbone/native-00004.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:a97f493300043c8b238f70fceaaab6b490f527ea2330cfabfe4a1c24aa23d185
3
+ size 3672318056
native_backbone/special_tokens_map.json ADDED
@@ -0,0 +1,63 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "additional_special_tokens": [
3
+ "<|im_start|>",
4
+ "<|im_end|>",
5
+ "<|action_start|>",
6
+ "<|action_end|>",
7
+ "<|interpreter|>",
8
+ "<|plugin|>",
9
+ "<restate>",
10
+ "</restate>",
11
+ "<planning>",
12
+ "</planning>",
13
+ "<recollect>",
14
+ "</recollect>",
15
+ "<execution>",
16
+ "</execution>",
17
+ "<review>",
18
+ "</review>",
19
+ "<summarize>",
20
+ "</summarize>",
21
+ "<retry>",
22
+ "</retry>",
23
+ "<conclude>",
24
+ "</conclude>",
25
+ "<img>",
26
+ "</img>",
27
+ "<IMG_CONTEXT>",
28
+ "<quad>",
29
+ "</quad>",
30
+ "<ref>",
31
+ "</ref>",
32
+ "<box>",
33
+ "</box>"
34
+ ],
35
+ "bos_token": {
36
+ "content": "<s>",
37
+ "lstrip": false,
38
+ "normalized": false,
39
+ "rstrip": false,
40
+ "single_word": false
41
+ },
42
+ "eos_token": {
43
+ "content": "<|im_end|>",
44
+ "lstrip": false,
45
+ "normalized": false,
46
+ "rstrip": false,
47
+ "single_word": false
48
+ },
49
+ "pad_token": {
50
+ "content": "</s>",
51
+ "lstrip": false,
52
+ "normalized": false,
53
+ "rstrip": false,
54
+ "single_word": false
55
+ },
56
+ "unk_token": {
57
+ "content": "<unk>",
58
+ "lstrip": false,
59
+ "normalized": false,
60
+ "rstrip": false,
61
+ "single_word": false
62
+ }
63
+ }
native_backbone/tokenization_internlm3.py ADDED
@@ -0,0 +1,294 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ from shutil import copyfile
3
+ from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple
4
+
5
+ import sentencepiece as spm
6
+ from transformers.tokenization_utils import AddedToken, PreTrainedTokenizer
7
+ from transformers.utils import logging
8
+
9
+ if TYPE_CHECKING:
10
+ from transformers.tokenization_utils_base import TextInput
11
+
12
+ logger = logging.get_logger(__name__)
13
+
14
+ VOCAB_FILES_NAMES = {"vocab_file": "tokenizer.model"}
15
+
16
+ SPIECE_UNDERLINE = "▁"
17
+
18
+
19
+ class InternLM3Tokenizer(PreTrainedTokenizer):
20
+ """
21
+ Construct a InternLM3 tokenizer. Based on byte-level Byte-Pair-Encoding. The default padding token is unset as there is
22
+ no padding token in the original model.
23
+
24
+ Args:
25
+ vocab_file (`str`):
26
+ Path to the vocabulary file.
27
+ unk_token (`str` or `tokenizers.AddedToken`, *optional*, defaults to `"<unk>"`):
28
+ The unknown token. A token that is not in the vocabulary cannot be converted to an ID and is set to be this
29
+ token instead.
30
+ bos_token (`str` or `tokenizers.AddedToken`, *optional*, defaults to `"<s>"`):
31
+ The beginning of sequence token that was used during pretraining. Can be used a sequence classifier token.
32
+ eos_token (`str` or `tokenizers.AddedToken`, *optional*, defaults to `"</s>"`):
33
+ The end of sequence token.
34
+ pad_token (`str` or `tokenizers.AddedToken`, *optional*):
35
+ A special token used to make arrays of tokens the same size for batching purpose. Will then be ignored by
36
+ attention mechanisms or loss computation.
37
+ sp_model_kwargs (`Dict[str, Any]`, `Optional`, *optional*):
38
+ Will be passed to the `SentencePieceProcessor.__init__()` method. The [Python wrapper for
39
+ SentencePiece](https://github.com/google/sentencepiece/tree/master/python) can be used, among other things,
40
+ to set:
41
+
42
+ - `enable_sampling`: Enable subword regularization.
43
+ - `nbest_size`: Sampling parameters for unigram. Invalid for BPE-Dropout.
44
+
45
+ - `nbest_size = {0,1}`: No sampling is performed.
46
+ - `nbest_size > 1`: samples from the nbest_size results.
47
+ - `nbest_size < 0`: assuming that nbest_size is infinite and samples from the all hypothesis (lattice)
48
+ using forward-filtering-and-backward-sampling algorithm.
49
+
50
+ - `alpha`: Smoothing parameter for unigram sampling, and dropout probability of merge operations for
51
+ BPE-dropout.
52
+
53
+ add_bos_token (`bool`, *optional*, defaults to `True`):
54
+ Whether or not to add an `bos_token` at the start of sequences.
55
+ add_eos_token (`bool`, *optional*, defaults to `False`):
56
+ Whether or not to add an `eos_token` at the end of sequences.
57
+ clean_up_tokenization_spaces (`bool`, *optional*, defaults to `False`):
58
+ Whether or not to cleanup spaces after decoding, cleanup consists in removing potential artifacts like
59
+ extra spaces.
60
+ use_default_system_prompt (`bool`, *optional*, defaults to `False`):
61
+ Whether or not the default system prompt for InternLM3 should be used.
62
+ spaces_between_special_tokens (`bool`, *optional*, defaults to `False`):
63
+ Whether or not to add spaces between special tokens.
64
+ spaces_for_interleaved_special_tokens (`bool`, *optional*, defaults to `False`):
65
+ Whether or not to add spaces between special tokens that are interleaved with normal tokens.
66
+ add_prefix_space (`bool`, *optional*, defaults to `True`):
67
+ Whether or not to add an initial space to the input. This allows to treat the leading word just as any
68
+ other word. Again, this should be set with `from_slow=True` to make sure it's taken into account.
69
+ """
70
+
71
+ vocab_files_names = VOCAB_FILES_NAMES
72
+ model_input_names = ["input_ids", "attention_mask"]
73
+
74
+ def __init__(
75
+ self,
76
+ vocab_file,
77
+ unk_token="<unk>",
78
+ bos_token="<s>",
79
+ eos_token="</s>",
80
+ pad_token=None,
81
+ sp_model_kwargs: Optional[Dict[str, Any]] = None,
82
+ add_bos_token=True,
83
+ add_eos_token=False,
84
+ clean_up_tokenization_spaces=False,
85
+ use_default_system_prompt=False,
86
+ spaces_between_special_tokens=False,
87
+ spaces_for_interleaved_special_tokens=False,
88
+ add_prefix_space=True,
89
+ **kwargs,
90
+ ):
91
+ self.sp_model_kwargs = {} if sp_model_kwargs is None else sp_model_kwargs
92
+ bos_token = AddedToken(bos_token, normalized=False, special=True) if isinstance(bos_token, str) else bos_token
93
+ eos_token = AddedToken(eos_token, normalized=False, special=True) if isinstance(eos_token, str) else eos_token
94
+ unk_token = AddedToken(unk_token, normalized=False, special=True) if isinstance(unk_token, str) else unk_token
95
+ pad_token = AddedToken(pad_token, normalized=False, special=True) if isinstance(pad_token, str) else pad_token
96
+
97
+ self.vocab_file = vocab_file
98
+ self.add_bos_token = add_bos_token
99
+ self.add_eos_token = add_eos_token
100
+ self.use_default_system_prompt = use_default_system_prompt
101
+ self.sp_model = spm.SentencePieceProcessor(**self.sp_model_kwargs)
102
+ self.sp_model.Load(vocab_file)
103
+ self.add_prefix_space = add_prefix_space
104
+ self.spaces_for_interleaved_special_tokens = spaces_for_interleaved_special_tokens
105
+
106
+ vocab_size = self.sp_model.get_piece_size()
107
+ self.decoder = {i: self.sp_model.id_to_piece(i) for i in range(vocab_size)}
108
+
109
+ super().__init__(
110
+ bos_token=bos_token,
111
+ eos_token=eos_token,
112
+ unk_token=unk_token,
113
+ pad_token=pad_token,
114
+ add_bos_token=add_bos_token,
115
+ add_eos_token=add_eos_token,
116
+ sp_model_kwargs=sp_model_kwargs,
117
+ clean_up_tokenization_spaces=clean_up_tokenization_spaces,
118
+ use_default_system_prompt=use_default_system_prompt,
119
+ spaces_between_special_tokens=spaces_between_special_tokens,
120
+ add_prefix_space=add_prefix_space,
121
+ **kwargs,
122
+ )
123
+
124
+ def __getstate__(self):
125
+ state = self.__dict__.copy()
126
+ state["sp_model"] = None
127
+ state["sp_model_proto"] = self.sp_model.serialized_model_proto()
128
+ return state
129
+
130
+ def __setstate__(self, d):
131
+ self.__dict__.update(d)
132
+ self.sp_model = spm.SentencePieceProcessor(**self.sp_model_kwargs)
133
+ self.sp_model.LoadFromSerializedProto(self.sp_model_proto)
134
+
135
+ @property
136
+ def vocab_size(self):
137
+ """Returns vocab size"""
138
+ return self.sp_model.get_piece_size()
139
+
140
+ def get_vocab(self):
141
+ """Returns vocab as a dict"""
142
+ vocab = {self.convert_ids_to_tokens(i): i for i in range(self.vocab_size)}
143
+ vocab.update(self.added_tokens_encoder)
144
+ return vocab
145
+
146
+ def tokenize(self, text: "TextInput", **kwargs) -> List[str]:
147
+ """
148
+ Args:
149
+ text: TextInput
150
+ Simply calls PreTrainedTokenizer's method
151
+ """
152
+ return super().tokenize(text, **kwargs)
153
+
154
+ def _tokenize(self, text, **kwargs):
155
+ """
156
+ Args:
157
+ text: TextInput
158
+ Returns a tokenized string. The Gemma tokenizer never adds a prefix space.
159
+ """
160
+ return self.sp_model.encode(text, out_type=str)
161
+
162
+ def _convert_token_to_id(self, token):
163
+ """Converts a token (str) in an id using the vocab."""
164
+ return self.sp_model.piece_to_id(token)
165
+
166
+ def _convert_id_to_token(self, index):
167
+ """Converts an index (integer) in a token (str) using the vocab."""
168
+ return self.decoder.get(index, "")
169
+
170
+ def convert_tokens_to_string(self, tokens):
171
+ """Converts a sequence of tokens (string) in a single string."""
172
+ # since we manually add the prefix space, we have to remove it when decoding
173
+ if tokens[0].startswith(SPIECE_UNDERLINE) and self.add_prefix_space:
174
+ tokens[0] = tokens[0][1:]
175
+
176
+ current_sub_tokens = []
177
+ out_string = ""
178
+ prev_is_special = False
179
+ for i, token in enumerate(tokens):
180
+ # make sure that special tokens are not decoded using sentencepiece model
181
+ if token in self.all_special_tokens:
182
+ if not prev_is_special and i != 0 and self.spaces_for_interleaved_special_tokens:
183
+ out_string += " "
184
+ out_string += self.sp_model.decode(current_sub_tokens) + token
185
+ prev_is_special = True
186
+ current_sub_tokens = []
187
+ else:
188
+ if (
189
+ prev_is_special
190
+ and i == 1
191
+ and self.add_prefix_space
192
+ and not token.startswith(SPIECE_UNDERLINE)
193
+ and self.spaces_for_interleaved_special_tokens
194
+ ):
195
+ out_string += " "
196
+ current_sub_tokens.append(token)
197
+ prev_is_special = False
198
+ out_string += self.sp_model.decode(current_sub_tokens)
199
+ return out_string
200
+
201
+ def save_vocabulary(self, save_directory, filename_prefix: Optional[str] = None) -> Tuple[str]:
202
+ """
203
+ Save the vocabulary and special tokens file to a directory.
204
+
205
+ Args:
206
+ save_directory (`str`):
207
+ The directory in which to save the vocabulary.
208
+
209
+ Returns:
210
+ `Tuple(str)`: Paths to the files saved.
211
+ """
212
+ if not os.path.isdir(save_directory):
213
+ logger.error(f"Vocabulary path ({save_directory}) should be a directory")
214
+ return
215
+ out_vocab_file = os.path.join(save_directory, (filename_prefix + "-" if filename_prefix else "") + VOCAB_FILES_NAMES["vocab_file"])
216
+
217
+ if os.path.abspath(self.vocab_file) != os.path.abspath(out_vocab_file) and os.path.isfile(self.vocab_file):
218
+ copyfile(self.vocab_file, out_vocab_file)
219
+ elif not os.path.isfile(self.vocab_file):
220
+ with open(out_vocab_file, "wb") as fi:
221
+ content_spiece_model = self.sp_model.serialized_model_proto()
222
+ fi.write(content_spiece_model)
223
+
224
+ return (out_vocab_file,)
225
+
226
+ def build_inputs_with_special_tokens(self, token_ids_0, token_ids_1=None):
227
+ bos_token_id = [self.bos_token_id] if self.add_bos_token else []
228
+ eos_token_id = [self.eos_token_id] if self.add_eos_token else []
229
+
230
+ output = bos_token_id + token_ids_0 + eos_token_id
231
+
232
+ if token_ids_1 is not None:
233
+ output = output + bos_token_id + token_ids_1 + eos_token_id
234
+
235
+ return output
236
+
237
+ def get_special_tokens_mask(
238
+ self, token_ids_0: List[int], token_ids_1: Optional[List[int]] = None, already_has_special_tokens: bool = False
239
+ ) -> List[int]:
240
+ """
241
+ Retrieve sequence ids from a token list that has no special tokens added. This method is called when adding
242
+ special tokens using the tokenizer `prepare_for_model` method.
243
+
244
+ Args:
245
+ token_ids_0 (`List[int]`):
246
+ List of IDs.
247
+ token_ids_1 (`List[int]`, *optional*):
248
+ Optional second list of IDs for sequence pairs.
249
+ already_has_special_tokens (`bool`, *optional*, defaults to `False`):
250
+ Whether or not the token list is already formatted with special tokens for the model.
251
+
252
+ Returns:
253
+ `List[int]`: A list of integers in the range [0, 1]: 1 for a special token, 0 for a sequence token.
254
+ """
255
+ if already_has_special_tokens:
256
+ return super().get_special_tokens_mask(token_ids_0=token_ids_0, token_ids_1=token_ids_1, already_has_special_tokens=True)
257
+
258
+ bos_token_id = [1] if self.add_bos_token else []
259
+ eos_token_id = [1] if self.add_eos_token else []
260
+
261
+ if token_ids_1 is None:
262
+ return bos_token_id + ([0] * len(token_ids_0)) + eos_token_id
263
+ return bos_token_id + ([0] * len(token_ids_0)) + eos_token_id + bos_token_id + ([0] * len(token_ids_1)) + eos_token_id
264
+
265
+ def create_token_type_ids_from_sequences(self, token_ids_0: List[int], token_ids_1: Optional[List[int]] = None) -> List[int]:
266
+ """
267
+ Creates a mask from the two sequences passed to be used in a sequence-pair classification task. An ALBERT
268
+ sequence pair mask has the following format:
269
+
270
+ ```
271
+ 0 0 0 0 0 0 0 0 0 0 0 1 1 1 1 1 1 1 1 1
272
+ | first sequence | second sequence |
273
+ ```
274
+
275
+ if token_ids_1 is None, only returns the first portion of the mask (0s).
276
+
277
+ Args:
278
+ token_ids_0 (`List[int]`):
279
+ List of ids.
280
+ token_ids_1 (`List[int]`, *optional*):
281
+ Optional second list of IDs for sequence pairs.
282
+
283
+ Returns:
284
+ `List[int]`: List of [token type IDs](../glossary#token-type-ids) according to the given sequence(s).
285
+ """
286
+ bos_token_id = [self.bos_token_id] if self.add_bos_token else []
287
+ eos_token_id = [self.eos_token_id] if self.add_eos_token else []
288
+
289
+ output = [0] * len(bos_token_id + token_ids_0 + eos_token_id)
290
+
291
+ if token_ids_1 is not None:
292
+ output += [1] * len(bos_token_id + token_ids_1 + eos_token_id)
293
+
294
+ return output
native_backbone/tokenizer.model ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:bcacff3229854f5103ee7a85473a30ca9a8b3a68f3aae9b7479574b23ac2256b
3
+ size 2475075
native_backbone/tokenizer_config.json ADDED
@@ -0,0 +1,330 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "add_bos_token": true,
3
+ "add_eos_token": false,
4
+ "add_prefix_space": true,
5
+ "added_tokens_decoder": {
6
+ "0": {
7
+ "content": "<unk>",
8
+ "lstrip": false,
9
+ "normalized": false,
10
+ "rstrip": false,
11
+ "single_word": false,
12
+ "special": true
13
+ },
14
+ "1": {
15
+ "content": "<s>",
16
+ "lstrip": false,
17
+ "normalized": false,
18
+ "rstrip": false,
19
+ "single_word": false,
20
+ "special": true
21
+ },
22
+ "2": {
23
+ "content": "</s>",
24
+ "lstrip": false,
25
+ "normalized": false,
26
+ "rstrip": false,
27
+ "single_word": false,
28
+ "special": true
29
+ },
30
+ "128111": {
31
+ "content": "<restate>",
32
+ "lstrip": false,
33
+ "normalized": false,
34
+ "rstrip": false,
35
+ "single_word": false,
36
+ "special": true
37
+ },
38
+ "128112": {
39
+ "content": "</restate>",
40
+ "lstrip": false,
41
+ "normalized": false,
42
+ "rstrip": false,
43
+ "single_word": false,
44
+ "special": true
45
+ },
46
+ "128113": {
47
+ "content": "<planning>",
48
+ "lstrip": false,
49
+ "normalized": false,
50
+ "rstrip": false,
51
+ "single_word": false,
52
+ "special": true
53
+ },
54
+ "128114": {
55
+ "content": "</planning>",
56
+ "lstrip": false,
57
+ "normalized": false,
58
+ "rstrip": false,
59
+ "single_word": false,
60
+ "special": true
61
+ },
62
+ "128115": {
63
+ "content": "<recollect>",
64
+ "lstrip": false,
65
+ "normalized": false,
66
+ "rstrip": false,
67
+ "single_word": false,
68
+ "special": true
69
+ },
70
+ "128116": {
71
+ "content": "</recollect>",
72
+ "lstrip": false,
73
+ "normalized": false,
74
+ "rstrip": false,
75
+ "single_word": false,
76
+ "special": true
77
+ },
78
+ "128117": {
79
+ "content": "<execution>",
80
+ "lstrip": false,
81
+ "normalized": false,
82
+ "rstrip": false,
83
+ "single_word": false,
84
+ "special": true
85
+ },
86
+ "128118": {
87
+ "content": "</execution>",
88
+ "lstrip": false,
89
+ "normalized": false,
90
+ "rstrip": false,
91
+ "single_word": false,
92
+ "special": true
93
+ },
94
+ "128119": {
95
+ "content": "<review>",
96
+ "lstrip": false,
97
+ "normalized": false,
98
+ "rstrip": false,
99
+ "single_word": false,
100
+ "special": true
101
+ },
102
+ "128120": {
103
+ "content": "</review>",
104
+ "lstrip": false,
105
+ "normalized": false,
106
+ "rstrip": false,
107
+ "single_word": false,
108
+ "special": true
109
+ },
110
+ "128121": {
111
+ "content": "<summarize>",
112
+ "lstrip": false,
113
+ "normalized": false,
114
+ "rstrip": false,
115
+ "single_word": false,
116
+ "special": true
117
+ },
118
+ "128122": {
119
+ "content": "</summarize>",
120
+ "lstrip": false,
121
+ "normalized": false,
122
+ "rstrip": false,
123
+ "single_word": false,
124
+ "special": true
125
+ },
126
+ "128123": {
127
+ "content": "<retry>",
128
+ "lstrip": false,
129
+ "normalized": false,
130
+ "rstrip": false,
131
+ "single_word": false,
132
+ "special": true
133
+ },
134
+ "128124": {
135
+ "content": "</retry>",
136
+ "lstrip": false,
137
+ "normalized": false,
138
+ "rstrip": false,
139
+ "single_word": false,
140
+ "special": true
141
+ },
142
+ "128125": {
143
+ "content": "<conclude>",
144
+ "lstrip": false,
145
+ "normalized": false,
146
+ "rstrip": false,
147
+ "single_word": false,
148
+ "special": true
149
+ },
150
+ "128126": {
151
+ "content": "</conclude>",
152
+ "lstrip": false,
153
+ "normalized": false,
154
+ "rstrip": false,
155
+ "single_word": false,
156
+ "special": true
157
+ },
158
+ "128127": {
159
+ "content": "<|plugin|>",
160
+ "lstrip": false,
161
+ "normalized": false,
162
+ "rstrip": false,
163
+ "single_word": false,
164
+ "special": true
165
+ },
166
+ "128128": {
167
+ "content": "<|interpreter|>",
168
+ "lstrip": false,
169
+ "normalized": false,
170
+ "rstrip": false,
171
+ "single_word": false,
172
+ "special": true
173
+ },
174
+ "128129": {
175
+ "content": "<|action_end|>",
176
+ "lstrip": false,
177
+ "normalized": false,
178
+ "rstrip": false,
179
+ "single_word": false,
180
+ "special": true
181
+ },
182
+ "128130": {
183
+ "content": "<|action_start|>",
184
+ "lstrip": false,
185
+ "normalized": false,
186
+ "rstrip": false,
187
+ "single_word": false,
188
+ "special": true
189
+ },
190
+ "128131": {
191
+ "content": "<|im_end|>",
192
+ "lstrip": false,
193
+ "normalized": false,
194
+ "rstrip": false,
195
+ "single_word": false,
196
+ "special": true
197
+ },
198
+ "128132": {
199
+ "content": "<|im_start|>",
200
+ "lstrip": false,
201
+ "normalized": false,
202
+ "rstrip": false,
203
+ "single_word": false,
204
+ "special": true
205
+ },
206
+ "128133": {
207
+ "content": "<img>",
208
+ "lstrip": false,
209
+ "normalized": false,
210
+ "rstrip": false,
211
+ "single_word": false,
212
+ "special": true
213
+ },
214
+ "128134": {
215
+ "content": "</img>",
216
+ "lstrip": false,
217
+ "normalized": false,
218
+ "rstrip": false,
219
+ "single_word": false,
220
+ "special": true
221
+ },
222
+ "128135": {
223
+ "content": "<IMG_CONTEXT>",
224
+ "lstrip": false,
225
+ "normalized": false,
226
+ "rstrip": false,
227
+ "single_word": false,
228
+ "special": true
229
+ },
230
+ "128136": {
231
+ "content": "<quad>",
232
+ "lstrip": false,
233
+ "normalized": false,
234
+ "rstrip": false,
235
+ "single_word": false,
236
+ "special": true
237
+ },
238
+ "128137": {
239
+ "content": "</quad>",
240
+ "lstrip": false,
241
+ "normalized": false,
242
+ "rstrip": false,
243
+ "single_word": false,
244
+ "special": true
245
+ },
246
+ "128138": {
247
+ "content": "<ref>",
248
+ "lstrip": false,
249
+ "normalized": false,
250
+ "rstrip": false,
251
+ "single_word": false,
252
+ "special": true
253
+ },
254
+ "128139": {
255
+ "content": "</ref>",
256
+ "lstrip": false,
257
+ "normalized": false,
258
+ "rstrip": false,
259
+ "single_word": false,
260
+ "special": true
261
+ },
262
+ "128140": {
263
+ "content": "<box>",
264
+ "lstrip": false,
265
+ "normalized": false,
266
+ "rstrip": false,
267
+ "single_word": false,
268
+ "special": true
269
+ },
270
+ "128141": {
271
+ "content": "</box>",
272
+ "lstrip": false,
273
+ "normalized": false,
274
+ "rstrip": false,
275
+ "single_word": false,
276
+ "special": true
277
+ }
278
+ },
279
+ "additional_special_tokens": [
280
+ "<|im_start|>",
281
+ "<|im_end|>",
282
+ "<|action_start|>",
283
+ "<|action_end|>",
284
+ "<|interpreter|>",
285
+ "<|plugin|>",
286
+ "<restate>",
287
+ "</restate>",
288
+ "<planning>",
289
+ "</planning>",
290
+ "<recollect>",
291
+ "</recollect>",
292
+ "<execution>",
293
+ "</execution>",
294
+ "<review>",
295
+ "</review>",
296
+ "<summarize>",
297
+ "</summarize>",
298
+ "<retry>",
299
+ "</retry>",
300
+ "<conclude>",
301
+ "</conclude>",
302
+ "<img>",
303
+ "</img>",
304
+ "<IMG_CONTEXT>",
305
+ "<quad>",
306
+ "</quad>",
307
+ "<ref>",
308
+ "</ref>",
309
+ "<box>",
310
+ "</box>"
311
+ ],
312
+ "auto_map": {
313
+ "AutoTokenizer": [
314
+ "tokenization_internlm3.InternLM3Tokenizer",
315
+ null
316
+ ]
317
+ },
318
+ "bos_token": "<s>",
319
+ "chat_template": "{{ bos_token }}{% for message in messages %}{{'<|im_start|>' + message['role'] + '\n' + message['content'] + '<|im_end|>' + '\n'}}{% endfor %}{% if add_generation_prompt %}{{ '<|im_start|>assistant\n' }}{% endif %}",
320
+ "clean_up_tokenization_spaces": false,
321
+ "eos_token": "<|im_end|>",
322
+ "extra_special_tokens": {},
323
+ "model_max_length": 8192,
324
+ "pad_token": "</s>",
325
+ "sp_model_kwargs": {},
326
+ "spaces_between_special_tokens": false,
327
+ "tokenizer_class": "InternLM3Tokenizer",
328
+ "unk_token": "<unk>",
329
+ "use_default_system_prompt": false
330
+ }
requirements.txt ADDED
@@ -0,0 +1,11 @@
 
 
 
 
 
 
 
 
 
 
 
 
1
+ torch==2.9.1
2
+ transformers==4.57.6
3
+ peft==0.17.1
4
+ timm
5
+ einops
6
+ accelerate
7
+ safetensors
8
+ huggingface_hub
9
+ Pillow
10
+ numpy
11
+ sentencepiece
source_gemma.py ADDED
@@ -0,0 +1,685 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Strict persistent-visual CVRR wrapper for Gemma 3 and Gemma 4.
2
+
3
+ The implementation intentionally wraps released Transformers models instead
4
+ of copying their decoder source. Native layer-call arguments are captured
5
+ during frozen prefix passes and replayed for the shared recurrent cell and
6
+ upper text-only continuation. This preserves each release's masks, rotary
7
+ geometry, and attention type while enforcing a visibly auditable interface:
8
+ the upper decoder receives only non-visual recurrent rows.
9
+ """
10
+
11
+ from __future__ import annotations
12
+
13
+ import contextlib
14
+ import io
15
+ import json
16
+ import math
17
+ import pathlib
18
+ from dataclasses import dataclass
19
+ from typing import Any
20
+
21
+ import torch
22
+ import torch.nn as nn
23
+ import torch.nn.functional as F
24
+
25
+ from .source_helpers import _layer_hidden
26
+
27
+
28
+ class LoRALinear(nn.Module):
29
+ def __init__(
30
+ self,
31
+ base: nn.Linear,
32
+ *,
33
+ rank: int,
34
+ alpha: float,
35
+ dropout: float,
36
+ ):
37
+ super().__init__()
38
+ if rank <= 0:
39
+ raise ValueError("LoRA rank must be positive")
40
+ self.base = base
41
+ for parameter in self.base.parameters():
42
+ parameter.requires_grad_(False)
43
+ self.rank = int(rank)
44
+ self.scale = float(alpha) / float(rank)
45
+ self.dropout = nn.Dropout(float(dropout))
46
+ self.lora_A = nn.Parameter(torch.empty(rank, base.in_features))
47
+ self.lora_B = nn.Parameter(torch.zeros(base.out_features, rank))
48
+ nn.init.kaiming_uniform_(self.lora_A, a=math.sqrt(5))
49
+ self.enabled = True
50
+
51
+ def forward(self, inputs):
52
+ result = self.base(inputs)
53
+ if not self.enabled:
54
+ return result
55
+ update = F.linear(self.dropout(inputs).float(), self.lora_A.float())
56
+ update = F.linear(update, self.lora_B.float())
57
+ return result + (update * self.scale).to(result.dtype)
58
+
59
+
60
+ def _module_parent(root: nn.Module, path: str):
61
+ parts = path.split(".")
62
+ parent = root
63
+ for part in parts[:-1]:
64
+ parent = getattr(parent, part)
65
+ return parent, parts[-1]
66
+
67
+
68
+ def inject_cell_lora(
69
+ cell: nn.Module,
70
+ *,
71
+ rank: int,
72
+ alpha: float,
73
+ dropout: float,
74
+ suffixes: set[str] | None = None,
75
+ ) -> dict[str, LoRALinear]:
76
+ if suffixes is None:
77
+ suffixes = {
78
+ "q_proj",
79
+ "k_proj",
80
+ "v_proj",
81
+ "o_proj",
82
+ "gate_proj",
83
+ "up_proj",
84
+ "down_proj",
85
+ }
86
+ selected = {
87
+ name: module
88
+ for name, module in cell.named_modules()
89
+ if isinstance(module, nn.Linear) and name.rsplit(".", 1)[-1] in suffixes
90
+ }
91
+ if not selected:
92
+ raise RuntimeError("no attention/MLP projections found in recurrent cell")
93
+ wrappers = {}
94
+ # Materialize the list before replacing children during traversal.
95
+ for name, module in selected.items():
96
+ parent, child = _module_parent(cell, name)
97
+ wrapper = LoRALinear(
98
+ module, rank=rank, alpha=alpha, dropout=dropout
99
+ )
100
+ setattr(parent, child, wrapper)
101
+ wrappers[name] = wrapper
102
+ return wrappers
103
+
104
+
105
+ @dataclass
106
+ class LayerCall:
107
+ positional_tail: tuple[Any, ...]
108
+ keywords: dict[str, Any]
109
+
110
+
111
+ @dataclass
112
+ class GemmaCVRRTrace:
113
+ """Frozen native states needed for one strict recurrent rollout.
114
+
115
+ ``scaffold`` is multimodal and is consumed only by the shared recurrent
116
+ cell. ``text_calls`` comes from an image-free sequence and is the sole
117
+ context replayed by the upper decoder. Keeping those two objects separate
118
+ makes the no-bypass contract directly inspectable.
119
+ """
120
+
121
+ scaffold: torch.Tensor
122
+ text_rows: torch.Tensor
123
+ visual_rows: torch.Tensor
124
+ question_valid: torch.Tensor
125
+ base_anchor: torch.Tensor
126
+ r1: torch.Tensor
127
+ mm_cell_call: LayerCall
128
+ text_calls: dict[int, LayerCall]
129
+
130
+ @property
131
+ def question_lengths(self) -> torch.Tensor:
132
+ return self.question_valid.sum(dim=1)
133
+
134
+ @property
135
+ def visual_lengths(self) -> torch.Tensor:
136
+ return self.visual_rows.sum(dim=1)
137
+
138
+
139
+ class _StopAfterCell(RuntimeError):
140
+ pass
141
+
142
+
143
+ def _capture_call(storage: dict[int, LayerCall], index: int):
144
+ def hook(_module, args, kwargs):
145
+ storage[index] = LayerCall(tuple(args[1:]), dict(kwargs))
146
+
147
+ return hook
148
+
149
+
150
+ def _gather_rows(full, mask):
151
+ batch, _, width = full.shape
152
+ lengths = mask.sum(dim=1).tolist()
153
+ result = full.new_zeros((batch, max(lengths), width))
154
+ for row, length in enumerate(lengths):
155
+ result[row, :length] = full[row, mask[row]]
156
+ return result
157
+
158
+
159
+ def _gather_ids(full, mask, *, pad_value: int):
160
+ lengths = mask.sum(dim=1).tolist()
161
+ result = full.new_full((full.shape[0], max(lengths)), pad_value)
162
+ valid = torch.zeros_like(result, dtype=torch.bool)
163
+ for row, length in enumerate(lengths):
164
+ result[row, :length] = full[row, mask[row]]
165
+ valid[row, :length] = True
166
+ return result, valid
167
+
168
+
169
+ def _replace_rows(full, mask, rows):
170
+ result = full.clone()
171
+ for batch_index in range(full.shape[0]):
172
+ count = int(mask[batch_index].sum())
173
+ result[batch_index, mask[batch_index]] = rows[batch_index, :count]
174
+ return result
175
+
176
+
177
+ class GemmaCVRR(nn.Module):
178
+ """One localized Gemma backbone with one shared recurrent-cell LoRA."""
179
+
180
+ def __init__(
181
+ self,
182
+ model_path: str,
183
+ *,
184
+ ell_star: int,
185
+ steps: int = 4,
186
+ beta: float = 0.33,
187
+ rank: int = 32,
188
+ alpha: float = 12.0,
189
+ dropout: float = 0.01,
190
+ device: str | torch.device = "cuda:0",
191
+ offline: bool = True,
192
+ ):
193
+ super().__init__()
194
+ from transformers import AutoConfig
195
+
196
+ self.model_path = str(model_path)
197
+ self.device_ref = torch.device(device)
198
+ config = AutoConfig.from_pretrained(
199
+ model_path, local_files_only=offline
200
+ )
201
+ if config.model_type == "gemma3":
202
+ from transformers import Gemma3ForConditionalGeneration as ModelClass
203
+ elif config.model_type == "gemma4_unified":
204
+ from transformers import (
205
+ Gemma4UnifiedForConditionalGeneration as ModelClass,
206
+ )
207
+ if int(getattr(config.text_config, "num_kv_shared_layers", 0)):
208
+ raise NotImplementedError(
209
+ "Gemma4 cross-layer shared KV would be an upper visual bypass"
210
+ )
211
+ else:
212
+ raise ValueError(f"unsupported Gemma model_type={config.model_type!r}")
213
+ self.base_model = ModelClass.from_pretrained(
214
+ model_path,
215
+ dtype=torch.bfloat16,
216
+ device_map=str(self.device_ref),
217
+ local_files_only=offline,
218
+ attn_implementation="sdpa",
219
+ )
220
+ self.model_type = config.model_type
221
+ self.layers = self.base_model.model.language_model.layers
222
+ self.ell_star = int(ell_star)
223
+ self.cell_index = self.ell_star + 1
224
+ self.upper_start = self.cell_index + 1
225
+ if not 0 <= self.ell_star <= len(self.layers) - 2:
226
+ raise ValueError(
227
+ f"ell_star={ell_star} must leave a recurrent cell and upper decoder"
228
+ )
229
+ if steps < 2:
230
+ raise ValueError("CVRR training requires at least two recurrent states")
231
+ if not 0.0 <= beta <= 1.0:
232
+ raise ValueError("beta must lie in [0,1]")
233
+ self.steps = int(steps)
234
+ self.beta = float(beta)
235
+ for parameter in self.base_model.parameters():
236
+ parameter.requires_grad_(False)
237
+ self.lora = inject_cell_lora(
238
+ self.layers[self.cell_index],
239
+ rank=rank,
240
+ alpha=alpha,
241
+ dropout=dropout,
242
+ )
243
+ # The dense checkpoint was placed before adapters were constructed;
244
+ # newly allocated A/B tensors otherwise remain on CPU until an outer
245
+ # training entrypoint happens to call ``model.to(device)``.
246
+ for module in self.lora.values():
247
+ module.to(self.device_ref)
248
+ self.rank = int(rank)
249
+ self.alpha = float(alpha)
250
+ self.adapter_dropout = float(dropout)
251
+ self.base_model.eval()
252
+
253
+ @contextlib.contextmanager
254
+ def adapters(self, enabled: bool):
255
+ previous = [module.enabled for module in self.lora.values()]
256
+ for module in self.lora.values():
257
+ module.enabled = bool(enabled)
258
+ try:
259
+ yield
260
+ finally:
261
+ for module, value in zip(self.lora.values(), previous):
262
+ module.enabled = value
263
+
264
+ def train(self, mode: bool = True):
265
+ super().train(mode)
266
+ # Frozen dense modules stay deterministic; only LoRA dropout follows
267
+ # training mode.
268
+ self.base_model.eval()
269
+ for module in self.lora.values():
270
+ module.dropout.train(mode)
271
+ return self
272
+
273
+ def _modality(self, mm_inputs):
274
+ if "token_type_ids" in mm_inputs:
275
+ return mm_inputs["token_type_ids"]
276
+ if "mm_token_type_ids" in mm_inputs:
277
+ return mm_inputs["mm_token_type_ids"]
278
+ raise ValueError("Gemma multimodal inputs have no modality IDs")
279
+
280
+ def _pad_token_id(self) -> int:
281
+ return int(self.base_model.config.text_config.pad_token_id)
282
+
283
+ def _initial_multimodal(self, mm_inputs):
284
+ calls: dict[int, LayerCall] = {}
285
+ captured = {}
286
+ cell = self.layers[self.cell_index]
287
+
288
+ def stop(_module, _args, output):
289
+ captured["hidden"] = _layer_hidden(output).detach()
290
+ raise _StopAfterCell
291
+
292
+ pre = cell.register_forward_pre_hook(
293
+ _capture_call(calls, self.cell_index), with_kwargs=True
294
+ )
295
+ post = cell.register_forward_hook(stop)
296
+ try:
297
+ with torch.no_grad(), self.adapters(False):
298
+ try:
299
+ self.base_model.model(
300
+ **mm_inputs, use_cache=False, return_dict=True
301
+ )
302
+ except _StopAfterCell:
303
+ pass
304
+ finally:
305
+ pre.remove()
306
+ post.remove()
307
+ if "hidden" not in captured or self.cell_index not in calls:
308
+ raise RuntimeError("failed to capture native multimodal recurrent cell")
309
+ return captured["hidden"], calls[self.cell_index]
310
+
311
+ def _text_context(self, question_ids, question_mask):
312
+ calls: dict[int, LayerCall] = {}
313
+ captured = {}
314
+ handles = []
315
+ for index in range(self.cell_index, len(self.layers)):
316
+ handles.append(
317
+ self.layers[index].register_forward_pre_hook(
318
+ _capture_call(calls, index), with_kwargs=True
319
+ )
320
+ )
321
+
322
+ def capture_cell(_module, _args, output):
323
+ captured["anchor"] = _layer_hidden(output).detach()
324
+
325
+ handles.append(self.layers[self.cell_index].register_forward_hook(capture_cell))
326
+ try:
327
+ with torch.no_grad(), self.adapters(False):
328
+ self.base_model.model(
329
+ input_ids=question_ids,
330
+ attention_mask=question_mask,
331
+ use_cache=False,
332
+ return_dict=True,
333
+ )
334
+ finally:
335
+ for handle in handles:
336
+ handle.remove()
337
+ missing = [
338
+ index
339
+ for index in range(self.cell_index, len(self.layers))
340
+ if index not in calls
341
+ ]
342
+ if missing or "anchor" not in captured:
343
+ raise RuntimeError(f"failed to capture text context; missing={missing}")
344
+ return captured["anchor"], calls
345
+
346
+ @staticmethod
347
+ def _call_layer(layer, hidden, call: LayerCall):
348
+ return _layer_hidden(
349
+ layer(hidden, *call.positional_tail, **call.keywords)
350
+ )
351
+
352
+ def _upper(self, state, text_calls):
353
+ hidden = state
354
+ with self.adapters(False):
355
+ for index in range(self.upper_start, len(self.layers)):
356
+ hidden = self._call_layer(
357
+ self.layers[index], hidden, text_calls[index]
358
+ )
359
+ hidden = self.base_model.model.language_model.norm(hidden)
360
+ logits = self.base_model.lm_head(hidden)
361
+ if self.model_type == "gemma4_unified":
362
+ cap = self.base_model.config.text_config.final_logit_softcapping
363
+ if cap is not None:
364
+ logits = torch.tanh(logits / cap) * cap
365
+ return logits
366
+
367
+ def extract(self, mm_inputs: dict[str, torch.Tensor]) -> GemmaCVRRTrace:
368
+ """Extract the frozen native first read and text-only upper context."""
369
+
370
+ attention = mm_inputs["attention_mask"].bool()
371
+ visual = self._modality(mm_inputs).eq(1) & attention
372
+ text_rows = (~visual) & attention
373
+ question_ids, question_valid = _gather_ids(
374
+ mm_inputs["input_ids"],
375
+ text_rows,
376
+ pad_value=self._pad_token_id(),
377
+ )
378
+ question_mask = question_valid.long()
379
+ first_full, mm_cell_call = self._initial_multimodal(mm_inputs)
380
+ base_anchor, text_calls = self._text_context(question_ids, question_mask)
381
+ r1 = _gather_rows(first_full, text_rows)
382
+ if r1.shape != base_anchor.shape:
383
+ raise RuntimeError(
384
+ "native multimodal and image-free question states are misaligned: "
385
+ f"R1={tuple(r1.shape)}, B={tuple(base_anchor.shape)}"
386
+ )
387
+ return GemmaCVRRTrace(
388
+ scaffold=first_full.detach(),
389
+ text_rows=text_rows,
390
+ visual_rows=visual,
391
+ question_valid=question_valid,
392
+ base_anchor=base_anchor.detach(),
393
+ r1=r1.detach(),
394
+ mm_cell_call=mm_cell_call,
395
+ text_calls=text_calls,
396
+ )
397
+
398
+ def rollout(
399
+ self,
400
+ trace: GemmaCVRRTrace,
401
+ *,
402
+ steps: int | None = None,
403
+ initial_state: torch.Tensor | None = None,
404
+ ) -> list[torch.Tensor]:
405
+ """Run the shared native cell and return ``[R1, ..., R_T]``."""
406
+
407
+ horizon = self.steps if steps is None else int(steps)
408
+ if horizon < 1:
409
+ raise ValueError("rollout steps must be positive")
410
+ state = trace.r1 if initial_state is None else initial_state
411
+ if state.shape != trace.r1.shape:
412
+ raise ValueError(
413
+ f"initial state shape {tuple(state.shape)} != {tuple(trace.r1.shape)}"
414
+ )
415
+ states = [state]
416
+ for _ in range(1, horizon):
417
+ recurrent_input = _replace_rows(
418
+ trace.scaffold, trace.text_rows, state
419
+ )
420
+ with self.adapters(True):
421
+ proposal_full = self._call_layer(
422
+ self.layers[self.cell_index],
423
+ recurrent_input,
424
+ trace.mm_cell_call,
425
+ )
426
+ proposal = _gather_rows(proposal_full, trace.text_rows)
427
+ state = state + self.beta * (proposal - state)
428
+ states.append(state)
429
+ return states
430
+
431
+ def decode_logits(
432
+ self,
433
+ trace: GemmaCVRRTrace,
434
+ state: torch.Tensor,
435
+ ) -> torch.Tensor:
436
+ """Decode one question-shaped state through the strict text-only path."""
437
+
438
+ if state.shape != trace.base_anchor.shape:
439
+ raise ValueError(
440
+ f"decoder state shape {tuple(state.shape)} != "
441
+ f"text anchor {tuple(trace.base_anchor.shape)}"
442
+ )
443
+ # Written explicitly as B + C_T to mirror the method definition. No
444
+ # multimodal row or multimodal cache is passed to `_upper`.
445
+ decoder_state = trace.base_anchor + (state - trace.base_anchor)
446
+ return self._upper(decoder_state, trace.text_calls).float()
447
+
448
+ def next_token_logits(
449
+ self,
450
+ trace: GemmaCVRRTrace,
451
+ state: torch.Tensor,
452
+ ) -> torch.Tensor:
453
+ """Return the distribution after each sample's final valid prompt row."""
454
+
455
+ logits = self.decode_logits(trace, state)
456
+ row = trace.question_lengths.to(logits.device) - 1
457
+ if bool((row < 0).any()):
458
+ raise ValueError("empty question sequence")
459
+ batch = torch.arange(logits.shape[0], device=logits.device)
460
+ return logits[batch, row]
461
+
462
+ def residual(self, trace: GemmaCVRRTrace, state: torch.Tensor) -> torch.Tensor:
463
+ return state - trace.base_anchor
464
+
465
+ def state_from_residual(
466
+ self,
467
+ trace: GemmaCVRRTrace,
468
+ residual: torch.Tensor,
469
+ ) -> torch.Tensor:
470
+ if residual.shape != trace.base_anchor.shape:
471
+ raise ValueError(
472
+ f"residual shape {tuple(residual.shape)} != "
473
+ f"text anchor {tuple(trace.base_anchor.shape)}"
474
+ )
475
+ return trace.base_anchor + residual
476
+
477
+ def forward(self, mm_inputs: dict[str, torch.Tensor], mm_labels):
478
+ question_labels, _ = _gather_ids(
479
+ mm_labels,
480
+ ((~self._modality(mm_inputs).eq(1)) & mm_inputs["attention_mask"].bool()),
481
+ pad_value=-100,
482
+ )
483
+ trace = self.extract(mm_inputs)
484
+ state = self.rollout(trace)[-1]
485
+ logits = self.decode_logits(trace, state)
486
+ shift_logits = logits[:, :-1]
487
+ shift_labels = question_labels[:, 1:]
488
+ token_loss = F.cross_entropy(
489
+ shift_logits.reshape(-1, shift_logits.shape[-1]),
490
+ shift_labels.reshape(-1),
491
+ ignore_index=-100,
492
+ reduction="none",
493
+ ).reshape(shift_labels.shape)
494
+ valid = shift_labels.ne(-100)
495
+ counts = valid.sum(dim=1).clamp_min(1)
496
+ per_example = (token_loss * valid).sum(dim=1) / counts
497
+ return {
498
+ "loss": per_example.mean(),
499
+ "logits": logits,
500
+ "labels": question_labels,
501
+ "r1": trace.r1.detach(),
502
+ "rT": state.detach(),
503
+ "base_anchor": trace.base_anchor.detach(),
504
+ "visual_rows": (
505
+ self._modality(mm_inputs).eq(1)
506
+ & mm_inputs["attention_mask"].bool()
507
+ ).sum(dim=1).detach(),
508
+ }
509
+
510
+ def adapter_state_dict(self):
511
+ return {
512
+ name: tensor.detach().cpu()
513
+ for name, tensor in self.state_dict().items()
514
+ if ".lora_A" in name or ".lora_B" in name
515
+ }
516
+
517
+ def save_adapter(self, output_dir: str | pathlib.Path, *, step: int):
518
+ output = pathlib.Path(output_dir)
519
+ output.mkdir(parents=True, exist_ok=True)
520
+ torch.save(self.adapter_state_dict(), output / "adapter_model.pt")
521
+ metadata = {
522
+ "format": "gemma_cvrr_lora_v1",
523
+ "base_model": self.model_path,
524
+ "model_type": self.model_type,
525
+ "ell_star": self.ell_star,
526
+ "cell_layer": self.cell_index,
527
+ "num_workspace_steps": self.steps,
528
+ "counterfactual_beta": self.beta,
529
+ "adapter_rank": self.rank,
530
+ "adapter_alpha": self.alpha,
531
+ "adapter_dropout": self.adapter_dropout,
532
+ "step": int(step),
533
+ "strict_path": True,
534
+ }
535
+ (output / "cvrr_config.json").write_text(json.dumps(metadata, indent=2))
536
+
537
+ def load_adapter(self, adapter_dir: str | pathlib.Path):
538
+ """Load an adapter exactly; missing or surplus LoRA tensors are fatal."""
539
+
540
+ adapter_dir = pathlib.Path(adapter_dir).expanduser().resolve()
541
+ metadata = json.loads((adapter_dir / "cvrr_config.json").read_text())
542
+ checks = {
543
+ "model_type": self.model_type,
544
+ "ell_star": self.ell_star,
545
+ "cell_layer": self.cell_index,
546
+ "num_workspace_steps": self.steps,
547
+ "adapter_rank": self.rank,
548
+ }
549
+ mismatches = {
550
+ name: (metadata.get(name), expected)
551
+ for name, expected in checks.items()
552
+ if metadata.get(name) != expected
553
+ }
554
+ if mismatches:
555
+ raise ValueError(f"adapter metadata mismatch: {mismatches}")
556
+ expected_keys = set(self.adapter_state_dict())
557
+ try:
558
+ payload = torch.load(
559
+ adapter_dir / "adapter_model.pt",
560
+ map_location="cpu",
561
+ weights_only=True,
562
+ )
563
+ except TypeError:
564
+ payload = torch.load(adapter_dir / "adapter_model.pt", map_location="cpu")
565
+ actual_keys = set(payload)
566
+ if actual_keys != expected_keys:
567
+ raise RuntimeError(
568
+ "adapter tensor mismatch: "
569
+ f"missing={sorted(expected_keys - actual_keys)[:8]}, "
570
+ f"unexpected={sorted(actual_keys - expected_keys)[:8]}"
571
+ )
572
+ incompatible = self.load_state_dict(payload, strict=False)
573
+ unexpected = list(incompatible.unexpected_keys)
574
+ missing_lora = [key for key in incompatible.missing_keys if key in expected_keys]
575
+ if unexpected or missing_lora:
576
+ raise RuntimeError(
577
+ f"adapter load failed: missing={missing_lora}, unexpected={unexpected}"
578
+ )
579
+ return metadata
580
+
581
+
582
+ def _prompt(processor, question: str, hint: str):
583
+ content = [
584
+ {"type": "image"},
585
+ {"type": "text", "text": question.strip() + str(hint)},
586
+ ]
587
+ return processor.apply_chat_template(
588
+ [{"role": "user", "content": content}],
589
+ tokenize=False,
590
+ add_generation_prompt=True,
591
+ )
592
+
593
+
594
+ class GemmaArrowCollator:
595
+ def __init__(self, processor):
596
+ self.processor = processor
597
+ self.tokenizer = processor.tokenizer
598
+
599
+ @staticmethod
600
+ def _pad_1d(values, pad):
601
+ width = max(item.shape[0] for item in values)
602
+ result = values[0].new_full((len(values), width), pad)
603
+ for index, item in enumerate(values):
604
+ result[index, : item.shape[0]] = item
605
+ return result
606
+
607
+ def __call__(self, features):
608
+ from PIL import Image
609
+
610
+ examples = []
611
+ for feature in features:
612
+ raw = feature["image_bytes"]
613
+ if isinstance(raw, memoryview):
614
+ raw = raw.tobytes()
615
+ with Image.open(io.BytesIO(raw)) as opened:
616
+ image = opened.convert("RGB")
617
+ prompt = _prompt(
618
+ self.processor,
619
+ str(feature["fixed_question"]),
620
+ str(feature["fixed_hint"]),
621
+ )
622
+ prompt_item = self.processor(
623
+ text=prompt, images=[image], return_tensors="pt"
624
+ )
625
+ full_item = self.processor(
626
+ text=prompt + str(feature["fixed_answer"]).strip(),
627
+ images=[image],
628
+ return_tensors="pt",
629
+ )
630
+ prompt_ids = prompt_item["input_ids"][0]
631
+ full_ids = full_item["input_ids"][0]
632
+ if not torch.equal(full_ids[: prompt_ids.numel()], prompt_ids):
633
+ raise RuntimeError("Gemma answer serialization changed the prompt prefix")
634
+ eos = full_ids.new_tensor([self.tokenizer.eos_token_id])
635
+ item = {name: value for name, value in full_item.items()}
636
+ item["input_ids"] = torch.cat((full_ids, eos))
637
+ item["attention_mask"] = torch.cat(
638
+ (item["attention_mask"][0], torch.ones_like(eos))
639
+ )
640
+ modality_name = (
641
+ "token_type_ids"
642
+ if "token_type_ids" in item
643
+ else "mm_token_type_ids"
644
+ )
645
+ item[modality_name] = torch.cat(
646
+ (item[modality_name][0], torch.zeros_like(eos))
647
+ )
648
+ answer_ids = item["input_ids"][prompt_ids.numel() :]
649
+ item["labels"] = torch.cat(
650
+ (torch.full_like(prompt_ids, -100), answer_ids)
651
+ )
652
+ examples.append(item)
653
+
654
+ sequence_names = {
655
+ "input_ids": int(self.tokenizer.pad_token_id),
656
+ "attention_mask": 0,
657
+ "labels": -100,
658
+ }
659
+ modality_name = (
660
+ "token_type_ids"
661
+ if "token_type_ids" in examples[0]
662
+ else "mm_token_type_ids"
663
+ )
664
+ sequence_names[modality_name] = 0
665
+ batch = {
666
+ name: self._pad_1d([item[name] for item in examples], pad)
667
+ for name, pad in sequence_names.items()
668
+ }
669
+ for name in examples[0]:
670
+ if name in sequence_names or name == "labels":
671
+ continue
672
+ values = [item[name] for item in examples]
673
+ batch[name] = torch.cat(values, dim=0)
674
+ labels = batch.pop("labels")
675
+ return {"mm_inputs": batch, "mm_labels": labels}
676
+
677
+
678
+ def move_batch(batch, device):
679
+ return {
680
+ "mm_inputs": {
681
+ name: value.to(device, non_blocking=True)
682
+ for name, value in batch["mm_inputs"].items()
683
+ },
684
+ "mm_labels": batch["mm_labels"].to(device, non_blocking=True),
685
+ }
source_helpers.py ADDED
@@ -0,0 +1,69 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ def _layer_hidden(output):
2
+ """Return the hidden tensor from a decoder-layer output."""
3
+
4
+ import torch
5
+
6
+ if torch.is_tensor(output):
7
+ return output
8
+ if isinstance(output, (tuple, list)) and output and torch.is_tensor(output[0]):
9
+ return output[0]
10
+ raise TypeError(f"unsupported decoder-layer output type: {type(output)!r}")
11
+
12
+ def _closest_ratio(aspect, ratios, width, height, image_size):
13
+ best = (1, 1)
14
+ difference = float("inf")
15
+ area = width * height
16
+ for ratio in ratios:
17
+ candidate = ratio[0] / ratio[1]
18
+ current = abs(aspect - candidate)
19
+ if current < difference or (
20
+ current == difference
21
+ and area > 0.5 * image_size * image_size * ratio[0] * ratio[1]
22
+ ):
23
+ difference = current
24
+ best = ratio
25
+ return best
26
+
27
+ def dynamic_tiles(image, *, image_size: int, max_tiles: int, thumbnail: bool):
28
+ """Official InternVL dynamic tiling, kept local for reproducibility."""
29
+
30
+ ratios = sorted(
31
+ {
32
+ (i, j)
33
+ for n in range(1, max_tiles + 1)
34
+ for i in range(1, n + 1)
35
+ for j in range(1, n + 1)
36
+ if 1 <= i * j <= max_tiles
37
+ },
38
+ key=lambda item: item[0] * item[1],
39
+ )
40
+ width, height = image.size
41
+ columns, rows = _closest_ratio(
42
+ width / height, ratios, width, height, image_size
43
+ )
44
+ resized = image.convert("RGB").resize(
45
+ (image_size * columns, image_size * rows), resample=3
46
+ )
47
+ tiles = []
48
+ for index in range(columns * rows):
49
+ left = (index % columns) * image_size
50
+ top = (index // columns) * image_size
51
+ tiles.append(
52
+ resized.crop((left, top, left + image_size, top + image_size))
53
+ )
54
+ if thumbnail and len(tiles) != 1:
55
+ tiles.append(image.convert("RGB").resize((image_size, image_size), 3))
56
+ return tiles
57
+
58
+ def _normalize_tiles(tiles):
59
+ import numpy as np
60
+ import torch
61
+
62
+ mean = torch.tensor((0.485, 0.456, 0.406)).view(3, 1, 1)
63
+ std = torch.tensor((0.229, 0.224, 0.225)).view(3, 1, 1)
64
+ tensors = []
65
+ for tile in tiles:
66
+ array = np.asarray(tile, dtype=np.float32) / 255.0
67
+ tensor = torch.from_numpy(array).permute(2, 0, 1)
68
+ tensors.append((tensor - mean) / std)
69
+ return torch.stack(tensors)
source_internvl.py ADDED
@@ -0,0 +1,286 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Strict persistent-visual CVRR wrapper for released InternVL3 chat models."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import copy
6
+ import io
7
+ import pathlib
8
+
9
+ import torch
10
+
11
+ from .source_gemma import (
12
+ GemmaCVRR,
13
+ LayerCall,
14
+ _StopAfterCell,
15
+ _capture_call,
16
+ _gather_ids,
17
+ _gather_rows,
18
+ _replace_rows,
19
+ inject_cell_lora,
20
+ )
21
+ from .source_helpers import _layer_hidden
22
+ from .source_helpers import _normalize_tiles, dynamic_tiles
23
+
24
+
25
+ class InternVLCVRR(GemmaCVRR):
26
+ def __init__(
27
+ self,
28
+ model_path: str,
29
+ *,
30
+ ell_star: int,
31
+ steps: int = 4,
32
+ beta: float = 0.33,
33
+ rank: int = 32,
34
+ alpha: float = 12.0,
35
+ dropout: float = 0.01,
36
+ device: str | torch.device = "cuda:0",
37
+ offline: bool = True,
38
+ ):
39
+ # Bypass GemmaCVRR.__init__, retaining its audited recurrence, adapter
40
+ # toggling, loss, and serialization methods.
41
+ torch.nn.Module.__init__(self)
42
+ from transformers import AutoModel, AutoTokenizer
43
+
44
+ import torch.distributed as dist
45
+
46
+ self._owns_process_group = False
47
+ if dist.is_available() and not dist.is_initialized():
48
+ import os
49
+ import tempfile
50
+
51
+ rendezvous = pathlib.Path(tempfile.gettempdir()) / f"cvrr_iv_train_{os.getpid()}"
52
+ dist.init_process_group(
53
+ "gloo", init_method=f"file://{rendezvous}", rank=0, world_size=1
54
+ )
55
+ self._owns_process_group = True
56
+
57
+ self.model_path = str(model_path)
58
+ self.device_ref = torch.device(device)
59
+ self.tokenizer = AutoTokenizer.from_pretrained(
60
+ model_path,
61
+ trust_remote_code=True,
62
+ use_fast=False,
63
+ local_files_only=offline,
64
+ )
65
+ self.base_model = AutoModel.from_pretrained(
66
+ model_path,
67
+ trust_remote_code=True,
68
+ local_files_only=offline,
69
+ low_cpu_mem_usage=True,
70
+ use_flash_attn=False,
71
+ dtype=torch.bfloat16,
72
+ device_map=str(self.device_ref),
73
+ )
74
+ self.model_type = "internvl_chat"
75
+ self.layers = self.base_model.language_model.model.layers
76
+ self.ell_star = int(ell_star)
77
+ self.cell_index = self.ell_star + 1
78
+ self.upper_start = self.cell_index + 1
79
+ if not 0 <= self.ell_star <= len(self.layers) - 2:
80
+ raise ValueError("ell_star must leave a cell and upper decoder")
81
+ if steps < 2 or not 0.0 <= beta <= 1.0:
82
+ raise ValueError("invalid recurrence depth or beta")
83
+ self.steps = int(steps)
84
+ self.beta = float(beta)
85
+ for parameter in self.base_model.parameters():
86
+ parameter.requires_grad_(False)
87
+ self.lora = inject_cell_lora(
88
+ self.layers[self.cell_index],
89
+ rank=rank,
90
+ alpha=alpha,
91
+ dropout=dropout,
92
+ suffixes={"wqkv", "wo", "w1", "w2", "w3"},
93
+ )
94
+ for module in self.lora.values():
95
+ module.to(self.device_ref)
96
+ self.rank = int(rank)
97
+ self.alpha = float(alpha)
98
+ self.adapter_dropout = float(dropout)
99
+ self.image_token_id = int(
100
+ self.tokenizer.convert_tokens_to_ids("<IMG_CONTEXT>")
101
+ )
102
+ self.base_model.img_context_token_id = self.image_token_id
103
+ self.base_model.eval()
104
+
105
+ def _pad_token_id(self) -> int:
106
+ return int(self.base_model.config.llm_config.pad_token_id)
107
+
108
+ def _modality(self, mm_inputs):
109
+ return mm_inputs["input_ids"].eq(self.image_token_id).long()
110
+
111
+ def _initial_multimodal(self, mm_inputs):
112
+ calls: dict[int, LayerCall] = {}
113
+ captured = {}
114
+ cell = self.layers[self.cell_index]
115
+
116
+ def stop(_module, _args, output):
117
+ captured["hidden"] = _layer_hidden(output).detach()
118
+ raise _StopAfterCell
119
+
120
+ pre = cell.register_forward_pre_hook(
121
+ _capture_call(calls, self.cell_index), with_kwargs=True
122
+ )
123
+ post = cell.register_forward_hook(stop)
124
+ try:
125
+ with torch.no_grad(), self.adapters(False):
126
+ try:
127
+ self.base_model(
128
+ **mm_inputs,
129
+ use_cache=False,
130
+ output_hidden_states=False,
131
+ return_dict=True,
132
+ )
133
+ except _StopAfterCell:
134
+ pass
135
+ finally:
136
+ pre.remove()
137
+ post.remove()
138
+ if "hidden" not in captured or self.cell_index not in calls:
139
+ raise RuntimeError("failed to capture InternVL recurrent cell")
140
+ return captured["hidden"], calls[self.cell_index]
141
+
142
+ def _text_context(self, question_ids, question_mask):
143
+ calls: dict[int, LayerCall] = {}
144
+ captured = {}
145
+ handles = []
146
+ for index in range(self.cell_index, len(self.layers)):
147
+ handles.append(
148
+ self.layers[index].register_forward_pre_hook(
149
+ _capture_call(calls, index), with_kwargs=True
150
+ )
151
+ )
152
+
153
+ def capture(_module, _args, output):
154
+ captured["anchor"] = _layer_hidden(output).detach()
155
+
156
+ handles.append(self.layers[self.cell_index].register_forward_hook(capture))
157
+ try:
158
+ with torch.no_grad(), self.adapters(False):
159
+ self.base_model.language_model(
160
+ input_ids=question_ids,
161
+ attention_mask=question_mask,
162
+ use_cache=False,
163
+ output_hidden_states=False,
164
+ return_dict=True,
165
+ )
166
+ finally:
167
+ for handle in handles:
168
+ handle.remove()
169
+ missing = [
170
+ index
171
+ for index in range(self.cell_index, len(self.layers))
172
+ if index not in calls
173
+ ]
174
+ if missing or "anchor" not in captured:
175
+ raise RuntimeError(f"failed to capture InternVL text path: {missing}")
176
+ return captured["anchor"], calls
177
+
178
+ def _upper(self, state, text_calls):
179
+ hidden = state
180
+ with self.adapters(False):
181
+ for index in range(self.upper_start, len(self.layers)):
182
+ hidden = self._call_layer(
183
+ self.layers[index], hidden, text_calls[index]
184
+ )
185
+ hidden = self.base_model.language_model.model.norm(hidden)
186
+ return self.base_model.language_model.output(hidden).float()
187
+
188
+
189
+ class InternVLArrowCollator:
190
+ def __init__(self, model: InternVLCVRR, *, max_tiles: int = 12):
191
+ self.tokenizer = model.tokenizer
192
+ self.template = copy.deepcopy(model.base_model.conv_template)
193
+ self.system_message = model.base_model.system_message
194
+ self.num_image_token = int(model.base_model.num_image_token)
195
+ self.image_size = int(
196
+ model.base_model.config.force_image_size
197
+ or model.base_model.config.vision_config.image_size
198
+ )
199
+ self.use_thumbnail = bool(model.base_model.config.use_thumbnail)
200
+ self.max_tiles = int(max_tiles)
201
+
202
+ @staticmethod
203
+ def _pad(values, pad):
204
+ width = max(value.shape[0] for value in values)
205
+ output = values[0].new_full((len(values), width), pad)
206
+ for index, value in enumerate(values):
207
+ output[index, : value.shape[0]] = value
208
+ return output
209
+
210
+ def _query(self, question, hint, num_tiles):
211
+ template = copy.deepcopy(self.template)
212
+ template.system_message = self.system_message
213
+ template.append_message(
214
+ template.roles[0],
215
+ "<image>\n" + str(question).strip() + str(hint),
216
+ )
217
+ template.append_message(template.roles[1], None)
218
+ query = template.get_prompt()
219
+ visual = (
220
+ "<img>"
221
+ + "<IMG_CONTEXT>" * self.num_image_token * num_tiles
222
+ + "</img>"
223
+ )
224
+ return query.replace("<image>", visual, 1)
225
+
226
+ def __call__(self, features):
227
+ from PIL import Image
228
+
229
+ rows = []
230
+ for feature in features:
231
+ raw = feature["image_bytes"]
232
+ if isinstance(raw, memoryview):
233
+ raw = raw.tobytes()
234
+ with Image.open(io.BytesIO(raw)) as opened:
235
+ tiles = dynamic_tiles(
236
+ opened.convert("RGB"),
237
+ image_size=self.image_size,
238
+ max_tiles=self.max_tiles,
239
+ thumbnail=self.use_thumbnail,
240
+ )
241
+ # The released InternViT does not cast inputs internally; its model
242
+ # card explicitly converts pixel_values to bfloat16 before forward.
243
+ pixels = _normalize_tiles(tiles).to(torch.bfloat16)
244
+ query = self._query(
245
+ feature["fixed_question"], feature["fixed_hint"], len(tiles)
246
+ )
247
+ tokenized = self.tokenizer(query, return_tensors="pt")
248
+ prompt = tokenized.input_ids[0]
249
+ full = self.tokenizer(
250
+ query + str(feature["fixed_answer"]).strip(),
251
+ return_tensors="pt",
252
+ ).input_ids[0]
253
+ if not torch.equal(full[: prompt.numel()], prompt):
254
+ raise RuntimeError("InternVL answer serialization changed the prompt prefix")
255
+ full = torch.cat(
256
+ (full, full.new_tensor([self.tokenizer.eos_token_id]))
257
+ )
258
+ answer = full[prompt.numel() :]
259
+ rows.append(
260
+ {
261
+ "input_ids": full,
262
+ "attention_mask": torch.cat(
263
+ (tokenized.attention_mask[0], torch.ones_like(answer))
264
+ ),
265
+ "labels": torch.cat((torch.full_like(prompt, -100), answer)),
266
+ "pixel_values": pixels,
267
+ "image_flags": torch.ones(len(tiles), 1, dtype=torch.long),
268
+ }
269
+ )
270
+ return {
271
+ "mm_inputs": {
272
+ "input_ids": self._pad(
273
+ [row["input_ids"] for row in rows], self.tokenizer.pad_token_id
274
+ ),
275
+ "attention_mask": self._pad(
276
+ [row["attention_mask"] for row in rows], 0
277
+ ),
278
+ "pixel_values": torch.cat(
279
+ [row["pixel_values"] for row in rows], dim=0
280
+ ),
281
+ "image_flags": torch.cat(
282
+ [row["image_flags"] for row in rows], dim=0
283
+ ),
284
+ },
285
+ "mm_labels": self._pad([row["labels"] for row in rows], -100),
286
+ }
source_perceive.py ADDED
@@ -0,0 +1,414 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Attention layouts for non-recurrent visual latent reasoning chains.
2
+
3
+ The sequence is laid out as::
4
+
5
+ [multimodal prompt ; clean question ; LOOK_1 ; THINK_1 ; ... ; answer]
6
+
7
+ All rows are processed once by the native VLM decoder. A block-sparse causal
8
+ graph makes LOOK rows the only latent rows with access to the multimodal
9
+ prefix, while THINK rows integrate one visual read without directly seeing
10
+ that prefix. Answer rows can see the clean question and THINK rows, but never
11
+ the multimodal prefix or LOOK rows. Consequently the image-dependent answer
12
+ path is
13
+
14
+ image -> LOOK_k -> THINK_k -> answer,
15
+
16
+ without a tied recurrent cell, a visual-state update, or an aggregation
17
+ module.
18
+
19
+ The transition-conditioned variant tightens the graph without adding a
20
+ module. The first LOOK/THINK pair bootstraps from the complete prompt and the
21
+ clean question. Every later pair sees the causal latent prefix, while its
22
+ LOOK row receives only visual placeholder rows from the multimodal prefix.
23
+ Thus later pairs cannot independently re-solve the original image/question
24
+ prompt; their task semantics must arrive through earlier latent states. The
25
+ last THINK is the sole latent exposed to answer rows and therefore acts as the
26
+ native-transformer aggregation state.
27
+ """
28
+
29
+ from __future__ import annotations
30
+
31
+ from dataclasses import dataclass
32
+
33
+ import torch
34
+
35
+
36
+ @dataclass(frozen=True)
37
+ class PerceiveDeliberateLayout:
38
+ """Physical row layout of one padded training/generation sequence."""
39
+
40
+ multimodal_length: int
41
+ question_length: int
42
+ num_pairs: int
43
+ answer_input_length: int = 0
44
+
45
+ def __post_init__(self) -> None:
46
+ if self.multimodal_length < 1:
47
+ raise ValueError("multimodal_length must be positive")
48
+ if self.question_length < 1:
49
+ raise ValueError("question_length must be positive")
50
+ if self.num_pairs < 1:
51
+ raise ValueError("num_pairs must be positive")
52
+ if self.answer_input_length < 0:
53
+ raise ValueError("answer_input_length must be non-negative")
54
+
55
+ @property
56
+ def num_latent_tokens(self) -> int:
57
+ return 2 * self.num_pairs
58
+
59
+ @property
60
+ def question_start(self) -> int:
61
+ return self.multimodal_length
62
+
63
+ @property
64
+ def latent_start(self) -> int:
65
+ return self.multimodal_length + self.question_length
66
+
67
+ @property
68
+ def answer_start(self) -> int:
69
+ return self.latent_start + self.num_latent_tokens
70
+
71
+ @property
72
+ def sequence_length(self) -> int:
73
+ return self.answer_start + self.answer_input_length
74
+
75
+ @property
76
+ def question_slice(self) -> slice:
77
+ return slice(self.question_start, self.latent_start)
78
+
79
+ @property
80
+ def answer_slice(self) -> slice:
81
+ return slice(self.answer_start, self.sequence_length)
82
+
83
+ @property
84
+ def look_indices(self) -> tuple[int, ...]:
85
+ return tuple(self.latent_start + 2 * index for index in range(self.num_pairs))
86
+
87
+ @property
88
+ def think_indices(self) -> tuple[int, ...]:
89
+ return tuple(index + 1 for index in self.look_indices)
90
+
91
+
92
+ def build_perceive_deliberate_mask(
93
+ layout: PerceiveDeliberateLayout,
94
+ multimodal_attention_mask: torch.Tensor,
95
+ question_attention_mask: torch.Tensor,
96
+ answer_input_attention_mask: torch.Tensor | None = None,
97
+ *,
98
+ multimodal_visual_mask: torch.Tensor | None = None,
99
+ local_visual_chain: bool = False,
100
+ transition_conditioned: bool = False,
101
+ question_visible_through_pair: int = 1,
102
+ disable_chain_links: bool = False,
103
+ disable_late_visual: bool = False,
104
+ final_pair_only: bool = False,
105
+ ) -> torch.BoolTensor:
106
+ """Build the strict LOOK/THINK/answer visibility graph.
107
+
108
+ The returned boolean mask has shape ``[B, 1, L, L]`` and follows PyTorch
109
+ SDPA semantics: ``True`` means that the query may attend to the key.
110
+
111
+ In the original graph, LOOK_k sees the valid multimodal prompt, the clean
112
+ question, and (for k>1) only THINK_{k-1}; THINK_k sees the clean question
113
+ and LOOK_k. In the transition-conditioned graph, pair 1 keeps those
114
+ bootstrap inputs, but every later latent sees its complete causal latent
115
+ prefix, later LOOK rows see only visual placeholder keys, and later THINK
116
+ rows normally receive no independent question shortcut. For a controlled
117
+ curriculum/intervention, ``question_visible_through_pair`` may temporarily
118
+ retain the clean-question edge through a later pair; ``1`` is the strict
119
+ inference graph. Answer rows see only the final THINK in that variant.
120
+ Every row also sees itself so padded query
121
+ rows remain numerically defined; padded rows are never exposed as keys to
122
+ semantic rows.
123
+ """
124
+
125
+ if multimodal_attention_mask.ndim != 2:
126
+ raise ValueError("multimodal_attention_mask must have shape [B, L_mm]")
127
+ if question_attention_mask.ndim != 2:
128
+ raise ValueError("question_attention_mask must have shape [B, L_q]")
129
+ batch_size = multimodal_attention_mask.shape[0]
130
+ expected_mm = (batch_size, layout.multimodal_length)
131
+ expected_q = (batch_size, layout.question_length)
132
+ if tuple(multimodal_attention_mask.shape) != expected_mm:
133
+ raise ValueError(
134
+ "multimodal mask/layout mismatch: "
135
+ f"expected {expected_mm}, got {tuple(multimodal_attention_mask.shape)}"
136
+ )
137
+ if tuple(question_attention_mask.shape) != expected_q:
138
+ raise ValueError(
139
+ "question mask/layout mismatch: "
140
+ f"expected {expected_q}, got {tuple(question_attention_mask.shape)}"
141
+ )
142
+ if layout.answer_input_length:
143
+ if answer_input_attention_mask is None:
144
+ raise ValueError("answer_input_attention_mask is required")
145
+ expected_answer = (batch_size, layout.answer_input_length)
146
+ if tuple(answer_input_attention_mask.shape) != expected_answer:
147
+ raise ValueError(
148
+ "answer mask/layout mismatch: "
149
+ f"expected {expected_answer}, got "
150
+ f"{tuple(answer_input_attention_mask.shape)}"
151
+ )
152
+ elif answer_input_attention_mask is not None and answer_input_attention_mask.numel():
153
+ raise ValueError("received answer mask for an empty answer-input segment")
154
+
155
+ device = multimodal_attention_mask.device
156
+ length = layout.sequence_length
157
+ visible = torch.zeros(
158
+ batch_size, length, length, dtype=torch.bool, device=device
159
+ )
160
+ mm_valid = multimodal_attention_mask.bool()
161
+ q_valid = question_attention_mask.bool()
162
+ visual_valid = None
163
+ if transition_conditioned or local_visual_chain:
164
+ if not 1 <= question_visible_through_pair <= layout.num_pairs:
165
+ raise ValueError(
166
+ "question_visible_through_pair must lie in "
167
+ f"[1,{layout.num_pairs}]"
168
+ )
169
+ if multimodal_visual_mask is None:
170
+ raise ValueError(
171
+ "multimodal_visual_mask is required for the "
172
+ "transition-conditioned graph"
173
+ )
174
+ if tuple(multimodal_visual_mask.shape) != expected_mm:
175
+ raise ValueError(
176
+ "visual mask/layout mismatch: "
177
+ f"expected {expected_mm}, got "
178
+ f"{tuple(multimodal_visual_mask.shape)}"
179
+ )
180
+ visual_valid = multimodal_visual_mask.bool()
181
+ if bool((visual_valid & ~mm_valid).any()):
182
+ raise ValueError("visual keys must be a subset of valid multimodal keys")
183
+
184
+ # The source multimodal prompt keeps the native causal graph. Padded query
185
+ # rows are harmless and receive a self edge below.
186
+ mm_causal = torch.ones(
187
+ layout.multimodal_length,
188
+ layout.multimodal_length,
189
+ dtype=torch.bool,
190
+ device=device,
191
+ ).tril()
192
+ visible[:, : layout.multimodal_length, : layout.multimodal_length] = (
193
+ mm_causal.unsqueeze(0) & mm_valid[:, None, :]
194
+ )
195
+
196
+ # The duplicated question is deliberately text-only: it is causally
197
+ # connected only to preceding valid rows of its own segment.
198
+ q_causal = torch.ones(
199
+ layout.question_length,
200
+ layout.question_length,
201
+ dtype=torch.bool,
202
+ device=device,
203
+ ).tril()
204
+ visible[
205
+ :, layout.question_slice, layout.question_slice
206
+ ] = q_causal.unsqueeze(0) & q_valid[:, None, :]
207
+
208
+ if local_visual_chain:
209
+ # Homogeneous one-pass latent chain. Visual memory is persistent and
210
+ # available at every position, while task state has exactly one local
211
+ # predecessor edge. No latent may consume the complete causal prefix.
212
+ latent_indices = range(layout.latent_start, layout.answer_start)
213
+ for latent_offset, latent in enumerate(latent_indices):
214
+ visible[:, latent, : layout.multimodal_length] = visual_valid
215
+ if latent_offset == 0:
216
+ visible[:, latent, layout.question_slice] = q_valid
217
+ else:
218
+ visible[:, latent, latent - 1] = True
219
+
220
+ if layout.answer_input_length:
221
+ answer_valid = answer_input_attention_mask.bool()
222
+ answer_causal = torch.ones(
223
+ layout.answer_input_length,
224
+ layout.answer_input_length,
225
+ dtype=torch.bool,
226
+ device=device,
227
+ ).tril()
228
+ visible[:, layout.answer_slice, layout.question_slice] = q_valid[:, None, :]
229
+ visible[:, layout.answer_slice, layout.answer_start - 1] = True
230
+ visible[:, layout.answer_slice, layout.answer_slice] = (
231
+ answer_causal.unsqueeze(0) & answer_valid[:, None, :]
232
+ )
233
+
234
+ diagonal = torch.arange(length, device=device)
235
+ visible[:, diagonal, diagonal] = True
236
+ return visible.unsqueeze(1)
237
+
238
+ for pair_index, (look, think) in enumerate(
239
+ zip(layout.look_indices, layout.think_indices)
240
+ ):
241
+ if final_pair_only and pair_index != layout.num_pairs - 1:
242
+ continue
243
+
244
+ if transition_conditioned and not final_pair_only:
245
+ # Pair 1 bootstraps the causal latent prefix. Later pairs cannot
246
+ # recover the task from the original prompt: only image placeholder
247
+ # rows remain visible outside the latent prefix.
248
+ if pair_index == 0:
249
+ visible[:, look, : layout.multimodal_length] = mm_valid
250
+ visible[:, look, layout.question_slice] = q_valid
251
+ visible[:, think, layout.question_slice] = q_valid
252
+ else:
253
+ if not disable_late_visual:
254
+ visible[:, look, : layout.multimodal_length] = visual_valid
255
+ if not disable_chain_links:
256
+ visible[:, look, layout.latent_start:look] = True
257
+ visible[:, think, layout.latent_start:look] = True
258
+ if pair_index < question_visible_through_pair:
259
+ visible[:, look, layout.question_slice] = q_valid
260
+ visible[:, think, layout.question_slice] = q_valid
261
+
262
+ # The local LOOK -> THINK edge is never an inter-pair chain-link
263
+ # intervention and therefore remains present in every arm.
264
+ visible[:, think, look] = True
265
+ else:
266
+ # Original graph, also used by the final-pair-only capacity
267
+ # control: each active pair may independently consume prompt + Q.
268
+ if not disable_late_visual or pair_index == 0 or final_pair_only:
269
+ visible[:, look, : layout.multimodal_length] = mm_valid
270
+ visible[:, look, layout.question_slice] = q_valid
271
+ if pair_index and not disable_chain_links and not final_pair_only:
272
+ visible[:, look, layout.think_indices[pair_index - 1]] = True
273
+
274
+ # Deliberation: no original multimodal key can be consumed here.
275
+ visible[:, think, layout.question_slice] = q_valid
276
+ visible[:, think, look] = True
277
+
278
+ if layout.answer_input_length:
279
+ answer_valid = answer_input_attention_mask.bool()
280
+ answer_causal = torch.ones(
281
+ layout.answer_input_length,
282
+ layout.answer_input_length,
283
+ dtype=torch.bool,
284
+ device=device,
285
+ ).tril()
286
+ visible[:, layout.answer_slice, layout.question_slice] = q_valid[:, None, :]
287
+ answer_thinks = (
288
+ [layout.think_indices[-1]]
289
+ if final_pair_only or transition_conditioned
290
+ else list(layout.think_indices)
291
+ )
292
+ visible[:, layout.answer_slice, answer_thinks] = True
293
+ visible[:, layout.answer_slice, layout.answer_slice] = (
294
+ answer_causal.unsqueeze(0) & answer_valid[:, None, :]
295
+ )
296
+
297
+ # Avoid all-masked softmax rows for physical padding. Semantic queries do
298
+ # not receive padded keys because every segment assignment above uses its
299
+ # validity mask.
300
+ diagonal = torch.arange(length, device=device)
301
+ visible[:, diagonal, diagonal] = True
302
+ return visible.unsqueeze(1)
303
+
304
+
305
+ def build_perceive_deliberate_positions(
306
+ layout: PerceiveDeliberateLayout,
307
+ multimodal_position_ids: torch.LongTensor,
308
+ multimodal_attention_mask: torch.Tensor,
309
+ question_attention_mask: torch.Tensor,
310
+ answer_input_attention_mask: torch.Tensor | None = None,
311
+ ) -> torch.LongTensor:
312
+ """Extend native multimodal M-RoPE with logical text positions.
313
+
314
+ The clean question keeps ordinary relative text positions but is shifted
315
+ after the largest valid multimodal M-RoPE coordinate. Latent and answer
316
+ rows then continue from each item's *logical* question length, independent
317
+ of right padding. All three M-RoPE coordinates are equal for these
318
+ non-spatial rows.
319
+ """
320
+
321
+ if multimodal_position_ids.ndim != 3 or multimodal_position_ids.shape[0] != 3:
322
+ raise ValueError("multimodal_position_ids must have shape [3, B, L_mm]")
323
+ batch_size = multimodal_attention_mask.shape[0]
324
+ if tuple(multimodal_position_ids.shape[1:]) != (
325
+ batch_size,
326
+ layout.multimodal_length,
327
+ ):
328
+ raise ValueError("multimodal position/layout mismatch")
329
+
330
+ mm_valid = multimodal_attention_mask.bool()
331
+ masked_mm = multimodal_position_ids.masked_fill(~mm_valid.unsqueeze(0), -1)
332
+ continuation = masked_mm.amax(dim=(0, 2)).clamp_min(-1) + 1 # [B]
333
+
334
+ q_valid = question_attention_mask.bool()
335
+ q_relative = q_valid.long().cumsum(dim=-1) - 1
336
+ q_relative = q_relative.masked_fill(~q_valid, 0)
337
+ q_positions = continuation[:, None] + q_relative
338
+ q_lengths = q_valid.sum(dim=-1)
339
+
340
+ latent_relative = torch.arange(
341
+ layout.num_latent_tokens,
342
+ device=multimodal_position_ids.device,
343
+ )
344
+ latent_positions = (
345
+ continuation[:, None] + q_lengths[:, None] + latent_relative[None, :]
346
+ )
347
+
348
+ segments = [multimodal_position_ids, q_positions.unsqueeze(0).expand(3, -1, -1)]
349
+ segments.append(latent_positions.unsqueeze(0).expand(3, -1, -1))
350
+ if layout.answer_input_length:
351
+ if answer_input_attention_mask is None:
352
+ raise ValueError("answer_input_attention_mask is required")
353
+ answer_relative = torch.arange(
354
+ layout.answer_input_length,
355
+ device=multimodal_position_ids.device,
356
+ )
357
+ answer_positions = (
358
+ continuation[:, None]
359
+ + q_lengths[:, None]
360
+ + layout.num_latent_tokens
361
+ + answer_relative[None, :]
362
+ )
363
+ segments.append(answer_positions.unsqueeze(0).expand(3, -1, -1))
364
+ return torch.cat(segments, dim=-1)
365
+
366
+
367
+ def build_answer_decode_mask(
368
+ layout: PerceiveDeliberateLayout,
369
+ question_attention_mask: torch.Tensor,
370
+ generated_attention_mask: torch.Tensor,
371
+ *,
372
+ local_visual_chain: bool = False,
373
+ final_pair_only: bool = False,
374
+ transition_conditioned: bool = False,
375
+ ) -> torch.BoolTensor:
376
+ """Visibility for one cached answer query after the latent prefill.
377
+
378
+ ``generated_attention_mask`` includes the current answer token. Physical
379
+ multimodal and LOOK cache rows remain present but are invisible.
380
+ """
381
+
382
+ if generated_attention_mask.ndim != 2:
383
+ raise ValueError("generated_attention_mask must have shape [B, N]")
384
+ batch_size, generated_length = generated_attention_mask.shape
385
+ if tuple(question_attention_mask.shape) != (
386
+ batch_size,
387
+ layout.question_length,
388
+ ):
389
+ raise ValueError("question mask/layout mismatch")
390
+ total_keys = layout.answer_start + generated_length
391
+ visible = torch.zeros(
392
+ batch_size, 1, 1, total_keys, dtype=torch.bool,
393
+ device=question_attention_mask.device,
394
+ )
395
+ visible[:, 0, 0, layout.question_slice] = question_attention_mask.bool()
396
+ if local_visual_chain:
397
+ visible[:, 0, 0, layout.answer_start - 1] = True
398
+ else:
399
+ answer_thinks = (
400
+ [layout.think_indices[-1]]
401
+ if final_pair_only or transition_conditioned
402
+ else list(layout.think_indices)
403
+ )
404
+ visible[:, 0, 0, answer_thinks] = True
405
+ visible[:, 0, 0, layout.answer_start:] = generated_attention_mask.bool()
406
+ return visible
407
+
408
+
409
+ __all__ = [
410
+ "PerceiveDeliberateLayout",
411
+ "build_answer_decode_mask",
412
+ "build_perceive_deliberate_mask",
413
+ "build_perceive_deliberate_positions",
414
+ ]
source_spatial.py ADDED
@@ -0,0 +1,349 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Question-conditioned recurrence over the complete native visual field.
2
+
3
+ The state is the dynamic-resolution sequence emitted by Qwen2.5-VL's vision
4
+ merger, before those image embeddings enter the language model. Tokens keep
5
+ their native two-dimensional grid throughout the recurrence. One shared cell
6
+ combines directional local messages with visual-to-question cross-attention::
7
+
8
+ V_0 = VisionMerger(image)
9
+ V_{k+1} = V_k + Cell(V_k, question, grid), k = 0, ..., T - 1
10
+
11
+ Only ``V_T`` is inserted into the ordinary Qwen multimodal prefix. The cell's
12
+ output projection is exactly zero at initialization, so every horizon,
13
+ including the default T=8, is initially bitwise identical to the base visual
14
+ embedding path. The full language model is then run from scratch; no cache
15
+ created while constructing the visual field is available to answer decoding.
16
+ """
17
+
18
+ from __future__ import annotations
19
+
20
+ from dataclasses import dataclass
21
+
22
+ import torch
23
+ import torch.nn as nn
24
+ import torch.nn.functional as F
25
+
26
+ from transformers.models.qwen2.modeling_qwen2 import Qwen2RMSNorm
27
+
28
+
29
+ @dataclass
30
+ class SpatialVisualTelemetry:
31
+ """Detached recurrence diagnostics; no visual trajectory is retained."""
32
+
33
+ update_rms: torch.Tensor # [T]
34
+ update_relative: torch.Tensor # [T]
35
+ step_drift: torch.Tensor # [T], 1 - cos(V_k, V_{k+1})
36
+ final_relative: torch.Tensor # scalar, RMS(V_T - V_0) / RMS(V_0)
37
+
38
+
39
+ class SpatialVisualRecurrentCell(nn.Module):
40
+ """Shared O(N) spatial cell for packed, variable-resolution image tokens.
41
+
42
+ Global N-by-N visual attention is deliberately absent: a sample may carry
43
+ 8192 merged visual tokens. Four directional grid messages preserve local
44
+ geometry in linear time, while low-rank visual-to-question attention makes
45
+ every patch update task dependent. The recurrent *state* remains at the
46
+ native language-model width; ``inner_width`` limits only the update
47
+ operator's compute and is not a state bottleneck.
48
+ """
49
+
50
+ def __init__(
51
+ self,
52
+ width: int,
53
+ inner_width: int,
54
+ num_heads: int,
55
+ *,
56
+ rms_norm_eps: float,
57
+ residual_scale: float,
58
+ ) -> None:
59
+ super().__init__()
60
+ if width < 1 or inner_width < 1:
61
+ raise ValueError("width and inner_width must be positive")
62
+ if num_heads < 1 or inner_width % num_heads:
63
+ raise ValueError(
64
+ "inner_width must be divisible by the positive num_heads"
65
+ )
66
+ if not 0.0 <= residual_scale <= 1.0:
67
+ raise ValueError("residual_scale must lie in [0, 1]")
68
+ self.width = int(width)
69
+ self.inner_width = int(inner_width)
70
+ self.num_heads = int(num_heads)
71
+ self.head_dim = self.inner_width // self.num_heads
72
+ self.residual_scale = float(residual_scale)
73
+
74
+ self.state_norm = Qwen2RMSNorm(self.width, eps=rms_norm_eps)
75
+ self.question_norm = Qwen2RMSNorm(self.width, eps=rms_norm_eps)
76
+
77
+ # One projection contains distinct north/south/west/east maps. This
78
+ # keeps orientation identifiable without allocating a dense 3x3 dxd
79
+ # convolution over the 3584-wide native visual state.
80
+ self.center_proj = nn.Linear(self.width, self.inner_width, bias=False)
81
+ self.neighbor_proj = nn.Linear(
82
+ self.width, 4 * self.inner_width, bias=False
83
+ )
84
+
85
+ self.query_proj = nn.Linear(self.width, self.inner_width, bias=False)
86
+ self.key_proj = nn.Linear(self.width, self.inner_width, bias=False)
87
+ self.value_proj = nn.Linear(self.width, self.inner_width, bias=False)
88
+ self.output_proj = nn.Linear(self.inner_width, self.width, bias=False)
89
+
90
+ def zero_output(self) -> None:
91
+ """Restore the exact base-model identity after generic HF init."""
92
+ with torch.no_grad():
93
+ self.output_proj.weight.zero_()
94
+
95
+ @staticmethod
96
+ def _token_layout(
97
+ merged_grid_thw: torch.LongTensor,
98
+ image_to_batch: torch.LongTensor,
99
+ batch_size: int,
100
+ expected_tokens: int,
101
+ ) -> tuple[torch.LongTensor, torch.LongTensor, torch.LongTensor]:
102
+ """Return token-to-batch ids, within-sample positions and counts."""
103
+ if merged_grid_thw.ndim != 2 or merged_grid_thw.shape[-1] != 3:
104
+ raise ValueError("merged_grid_thw must have shape [num_images, 3]")
105
+ if image_to_batch.ndim != 1 or image_to_batch.shape[0] != len(
106
+ merged_grid_thw
107
+ ):
108
+ raise ValueError("image_to_batch must contain one id per image")
109
+ per_image = merged_grid_thw.prod(dim=-1).long()
110
+ if int(per_image.sum()) != int(expected_tokens):
111
+ raise ValueError(
112
+ "merged grids and visual field disagree: "
113
+ f"grid tokens={int(per_image.sum())}, state tokens={expected_tokens}"
114
+ )
115
+ token_to_batch = torch.repeat_interleave(image_to_batch, per_image)
116
+ if token_to_batch.numel() and (
117
+ int(token_to_batch.min()) < 0
118
+ or int(token_to_batch.max()) >= batch_size
119
+ ):
120
+ raise ValueError("image_to_batch contains an out-of-range batch id")
121
+ if token_to_batch.numel() > 1 and bool(
122
+ (token_to_batch[1:] < token_to_batch[:-1]).any()
123
+ ):
124
+ raise ValueError(
125
+ "images must follow the same batch-major order as Qwen inputs"
126
+ )
127
+ counts = torch.bincount(token_to_batch, minlength=batch_size)
128
+ starts = counts.cumsum(0) - counts
129
+ within = torch.arange(
130
+ expected_tokens, device=token_to_batch.device
131
+ ) - torch.repeat_interleave(starts, counts)
132
+ return token_to_batch, within, counts
133
+
134
+ @staticmethod
135
+ def _neighbor_edges(
136
+ merged_grid_thw: torch.LongTensor,
137
+ expected_tokens: int,
138
+ ) -> tuple[torch.LongTensor, torch.LongTensor, torch.LongTensor]:
139
+ """Build packed directed 4-neighbour edges once for all recurrent steps.
140
+
141
+ Direction ids are 0/1/2/3 = north/south/west/east *as seen by the
142
+ destination token*. Temporal slices are separate 2-D fields.
143
+ """
144
+ device = merged_grid_thw.device
145
+ destinations: list[torch.Tensor] = []
146
+ sources: list[torch.Tensor] = []
147
+ directions: list[torch.Tensor] = []
148
+ offset = 0
149
+ for t_value, h_value, w_value in merged_grid_thw.detach().cpu().tolist():
150
+ t, h, w = int(t_value), int(h_value), int(w_value)
151
+ count = t * h * w
152
+ if min(t, h, w) < 1:
153
+ raise ValueError(f"invalid merged visual grid {(t, h, w)}")
154
+ index = torch.arange(
155
+ offset, offset + count, device=device, dtype=torch.long
156
+ ).view(t, h, w)
157
+
158
+ def append(dst: torch.Tensor, src: torch.Tensor, direction: int) -> None:
159
+ if dst.numel() == 0:
160
+ return
161
+ destinations.append(dst.reshape(-1))
162
+ sources.append(src.reshape(-1))
163
+ directions.append(
164
+ torch.full(
165
+ (dst.numel(),), direction, device=device, dtype=torch.long
166
+ )
167
+ )
168
+
169
+ append(index[:, 1:, :], index[:, :-1, :], 0) # north
170
+ append(index[:, :-1, :], index[:, 1:, :], 1) # south
171
+ append(index[:, :, 1:], index[:, :, :-1], 2) # west
172
+ append(index[:, :, :-1], index[:, :, 1:], 3) # east
173
+ offset += count
174
+ if offset != expected_tokens:
175
+ raise ValueError(
176
+ f"edge grids contain {offset} tokens, expected {expected_tokens}"
177
+ )
178
+ if not destinations:
179
+ empty = torch.empty(0, device=device, dtype=torch.long)
180
+ return empty, empty, empty
181
+ return (
182
+ torch.cat(destinations),
183
+ torch.cat(sources),
184
+ torch.cat(directions),
185
+ )
186
+
187
+ def _question_attention(
188
+ self,
189
+ normalized_state: torch.Tensor,
190
+ question_keys: torch.Tensor,
191
+ question_values: torch.Tensor,
192
+ question_valid: torch.BoolTensor,
193
+ token_to_batch: torch.LongTensor,
194
+ within_sample: torch.LongTensor,
195
+ token_counts: torch.LongTensor,
196
+ ) -> torch.Tensor:
197
+ """Visual-query/text-memory SDPA for a packed visual sequence."""
198
+ batch_size = question_keys.shape[0]
199
+ max_tokens = int(token_counts.max()) if token_counts.numel() else 0
200
+ if max_tokens == 0:
201
+ return normalized_state.new_empty(0, self.inner_width)
202
+
203
+ query_flat = self.query_proj(normalized_state)
204
+ query = query_flat.new_zeros(
205
+ batch_size, max_tokens, self.inner_width
206
+ )
207
+ query[token_to_batch, within_sample] = query_flat
208
+
209
+ bsz, question_length, _ = question_keys.shape
210
+ query = query.view(
211
+ bsz, max_tokens, self.num_heads, self.head_dim
212
+ ).transpose(1, 2)
213
+ key = question_keys.view(
214
+ bsz, question_length, self.num_heads, self.head_dim
215
+ ).transpose(1, 2)
216
+ value = question_values.view(
217
+ bsz, question_length, self.num_heads, self.head_dim
218
+ ).transpose(1, 2)
219
+ # Boolean SDPA masks use True for entries that are allowed to attend.
220
+ allowed = question_valid[:, None, None, :]
221
+ attended = F.scaled_dot_product_attention(
222
+ query,
223
+ key,
224
+ value,
225
+ attn_mask=allowed,
226
+ dropout_p=0.0,
227
+ is_causal=False,
228
+ )
229
+ attended = attended.transpose(1, 2).reshape(
230
+ bsz, max_tokens, self.inner_width
231
+ )
232
+ return attended[token_to_batch, within_sample]
233
+
234
+ def forward(
235
+ self,
236
+ visual_state: torch.Tensor,
237
+ merged_grid_thw: torch.LongTensor,
238
+ image_to_batch: torch.LongTensor,
239
+ question_embeddings: torch.Tensor,
240
+ question_attention_mask: torch.Tensor | None,
241
+ *,
242
+ steps: int,
243
+ ) -> tuple[torch.Tensor, SpatialVisualTelemetry]:
244
+ if visual_state.ndim != 2 or visual_state.shape[-1] != self.width:
245
+ raise ValueError(
246
+ f"visual_state must be [N,{self.width}], got "
247
+ f"{tuple(visual_state.shape)}"
248
+ )
249
+ if question_embeddings.ndim != 3 or question_embeddings.shape[-1] != self.width:
250
+ raise ValueError(
251
+ f"question_embeddings must be [B,Q,{self.width}]"
252
+ )
253
+ if steps < 0:
254
+ raise ValueError("steps must be non-negative")
255
+ batch_size = question_embeddings.shape[0]
256
+ token_to_batch, within_sample, token_counts = self._token_layout(
257
+ merged_grid_thw,
258
+ image_to_batch,
259
+ batch_size,
260
+ visual_state.shape[0],
261
+ )
262
+ edge_dst, edge_src, edge_direction = self._neighbor_edges(
263
+ merged_grid_thw, visual_state.shape[0]
264
+ )
265
+
266
+ if question_attention_mask is None:
267
+ question_valid = torch.ones(
268
+ question_embeddings.shape[:2],
269
+ dtype=torch.bool,
270
+ device=question_embeddings.device,
271
+ )
272
+ else:
273
+ if question_attention_mask.shape != question_embeddings.shape[:2]:
274
+ raise ValueError("question attention mask shape mismatch")
275
+ question_valid = question_attention_mask > 0
276
+ if not bool(question_valid.any(dim=-1).all()):
277
+ raise ValueError("every image-bearing sample needs a question token")
278
+
279
+ normalized_question = self.question_norm(question_embeddings)
280
+ question_keys = self.key_proj(normalized_question)
281
+ question_values = self.value_proj(normalized_question)
282
+
283
+ initial = visual_state
284
+ state = visual_state
285
+ update_rms: list[torch.Tensor] = []
286
+ update_relative: list[torch.Tensor] = []
287
+ step_drift: list[torch.Tensor] = []
288
+ for _ in range(steps):
289
+ normalized = self.state_norm(state)
290
+ center = self.center_proj(normalized)
291
+ directional = self.neighbor_proj(normalized).view(
292
+ state.shape[0], 4, self.inner_width
293
+ )
294
+ spatial = center.new_zeros(center.shape)
295
+ degree = center.new_zeros(center.shape[0], 1)
296
+ if edge_dst.numel():
297
+ messages = directional[edge_src, edge_direction]
298
+ spatial.index_add_(0, edge_dst, messages)
299
+ degree.index_add_(
300
+ 0,
301
+ edge_dst,
302
+ torch.ones(
303
+ edge_dst.shape[0], 1, dtype=center.dtype, device=center.device
304
+ ),
305
+ )
306
+ spatial = spatial / degree.clamp_min(1.0)
307
+ question = self._question_attention(
308
+ normalized,
309
+ question_keys,
310
+ question_values,
311
+ question_valid,
312
+ token_to_batch,
313
+ within_sample,
314
+ token_counts,
315
+ )
316
+ mixed = F.silu((center + spatial + question) / (3.0**0.5))
317
+ update = self.output_proj(mixed) * self.residual_scale
318
+ next_state = state + update
319
+
320
+ with torch.no_grad():
321
+ state_rms = state.float().square().mean().sqrt().clamp_min(1e-8)
322
+ update_norm = update.float().square().mean().sqrt()
323
+ cosine = F.cosine_similarity(
324
+ state.float(), next_state.float(), dim=-1
325
+ ).mean()
326
+ update_rms.append(update_norm.detach())
327
+ update_relative.append((update_norm / state_rms).detach())
328
+ step_drift.append((1.0 - cosine).detach())
329
+ state = next_state
330
+
331
+ with torch.no_grad():
332
+ initial_rms = initial.float().square().mean().sqrt().clamp_min(1e-8)
333
+ final_relative = (
334
+ (state.float() - initial.float()).square().mean().sqrt()
335
+ / initial_rms
336
+ ).detach()
337
+ empty = initial.new_empty(0, dtype=torch.float32)
338
+ telemetry = SpatialVisualTelemetry(
339
+ update_rms=(torch.stack(update_rms) if update_rms else empty),
340
+ update_relative=(
341
+ torch.stack(update_relative) if update_relative else empty
342
+ ),
343
+ step_drift=(torch.stack(step_drift) if step_drift else empty),
344
+ final_relative=final_relative,
345
+ )
346
+ return state, telemetry
347
+
348
+
349
+ __all__ = ["SpatialVisualRecurrentCell", "SpatialVisualTelemetry"]
source_splitting.py ADDED
@@ -0,0 +1,473 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Layer-split forward for Qwen2.5-VL and Qwen3-VL.
2
+
3
+ Method §3.2 needs ``F = F_{>l*} o F_{<=l*}`` as two separately runnable halves so
4
+ that the multimodal branch can be cut off at ``l*`` and replaced by the
5
+ workspace. §3.1 needs the same split to read and patch activations at a chosen
6
+ depth.
7
+
8
+ This module reimplements the prologue of the native Qwen VL text-model forward
9
+ -- embedding merge, M-RoPE index, causal mask, rotary embeddings -- as a
10
+ reusable :class:`SplitContext`, then exposes the decoder-layer loop as a range
11
+ you can run piecewise. Qwen3-VL additionally injects three ``DeepStack``
12
+ vision features after language layers 0--2; those tensors are carried in the
13
+ split context and applied at the identical layer boundaries. It deliberately
14
+ mirrors transformers 4.57.6 rather than monkeypatching it;
15
+ ``tests/test_split_equivalence.py`` asserts the composed halves reproduce the
16
+ stock forward bit-for-bit, which is what makes the mirroring safe to rely on.
17
+
18
+ Shapes use ``B`` batch, ``L`` sequence, ``d`` backbone width (3584 on the 7B),
19
+ ``N_v`` visual tokens, ``N_q`` question tokens.
20
+ """
21
+
22
+ from __future__ import annotations
23
+
24
+ from dataclasses import dataclass, replace
25
+ from typing import Any
26
+
27
+ import torch
28
+ from transformers.cache_utils import Cache
29
+ from transformers.masking_utils import (
30
+ create_causal_mask,
31
+ create_sliding_window_causal_mask,
32
+ )
33
+
34
+
35
+ @dataclass
36
+ class SplitContext:
37
+ """Per-forward state shared by every decoder layer.
38
+
39
+ Computed once by :func:`make_split_context` so that layer ranges can be run
40
+ independently without recomputing masks or rotary tables.
41
+ """
42
+
43
+ hidden_states: torch.Tensor # [B, L, d] - mutated as layers run
44
+ position_ids: torch.Tensor # [3, B, L] - M-RoPE (t, h, w)
45
+ position_embeddings: tuple[torch.Tensor, torch.Tensor] # (cos, sin) [B, L, head_dim]
46
+ causal_mask_mapping: dict[str, torch.Tensor | None]
47
+ cache_position: torch.Tensor # [L]
48
+ text_position_ids: torch.Tensor | None # [B, L] only when packed
49
+ past_key_values: Cache | None
50
+ # Original 2-D key-padding mask. ``create_causal_mask`` is allowed to
51
+ # return ``None`` for SDPA and delegate causality to ``is_causal``; keeping
52
+ # this tensor lets counterfactual branches materialize the equivalent mask
53
+ # before removing a precisely selected set of attention edges.
54
+ attention_mask: torch.Tensor | None = None # [B, L_kv]
55
+ # Qwen3-VL only. DeepStack adds one visual feature tensor after each of
56
+ # the first three language layers. They stay ``None`` for Qwen2.5-VL and
57
+ # for every text-only branch.
58
+ visual_pos_masks: torch.Tensor | None = None # [B, L] bool
59
+ deepstack_visual_embeds: list[torch.Tensor] | None = None
60
+
61
+ def clone_at(self, hidden_states: torch.Tensor) -> "SplitContext":
62
+ """Same context, different hidden states (for patched re-runs)."""
63
+ return SplitContext(
64
+ hidden_states=hidden_states,
65
+ position_ids=self.position_ids,
66
+ position_embeddings=self.position_embeddings,
67
+ causal_mask_mapping=self.causal_mask_mapping,
68
+ cache_position=self.cache_position,
69
+ text_position_ids=self.text_position_ids,
70
+ past_key_values=self.past_key_values,
71
+ attention_mask=self.attention_mask,
72
+ visual_pos_masks=self.visual_pos_masks,
73
+ deepstack_visual_embeds=self.deepstack_visual_embeds,
74
+ )
75
+
76
+
77
+ # ---------------------------------------------------------------------------
78
+ # embedding / position construction
79
+ # ---------------------------------------------------------------------------
80
+
81
+
82
+ def embed_multimodal(
83
+ vl_model,
84
+ input_ids: torch.LongTensor, # [B, L]
85
+ pixel_values: torch.Tensor | None = None,
86
+ image_grid_thw: torch.LongTensor | None = None,
87
+ attention_mask: torch.Tensor | None = None,
88
+ *,
89
+ return_deepstack: bool = False,
90
+ ) -> (
91
+ tuple[torch.Tensor, torch.Tensor]
92
+ | tuple[
93
+ torch.Tensor,
94
+ torch.Tensor,
95
+ torch.Tensor | None,
96
+ list[torch.Tensor] | None,
97
+ ]
98
+ ):
99
+ """Token embeddings with image features scattered in, plus M-RoPE indices.
100
+
101
+ Mirrors the prefill path of ``Qwen2_5_VLModel.forward``. Pass
102
+ ``pixel_values=None`` to get the text-only branch used for ``Q*``.
103
+
104
+ Args:
105
+ vl_model: a native ``Qwen2_5_VLModel`` or ``Qwen3VLModel`` (i.e.
106
+ ``model.model``, not the ``...ForConditionalGeneration`` wrapper).
107
+ return_deepstack: also return Qwen3-VL's visual-position mask and
108
+ DeepStack features. The default two-value return keeps all
109
+ Qwen2.5 callers backward compatible.
110
+
111
+ Returns:
112
+ ``(inputs_embeds [B, L, d], position_ids [3, B, L])`` and optionally
113
+ ``(visual_pos_masks, deepstack_visual_embeds)``.
114
+ """
115
+ inputs_embeds = vl_model.get_input_embeddings()(input_ids) # [B, L, d]
116
+ model_type = str(getattr(vl_model.config, "model_type", ""))
117
+ # A freshly wrapped model exposes the native qwen3_vl config here.
118
+ # Reloading a fully saved CLOSE checkpoint reconstructs the nested native
119
+ # backbone from the wrapper config, so its model view carries
120
+ # close_qwen3_vl instead. Both use the tuple-returning Qwen3 feature API.
121
+ is_qwen3_vl = model_type in {"qwen3_vl", "close_qwen3_vl"}
122
+ visual_pos_masks = None
123
+ deepstack_visual_embeds = None
124
+
125
+ if pixel_values is not None:
126
+ image_features = vl_model.get_image_features(pixel_values, image_grid_thw)
127
+ if is_qwen3_vl:
128
+ image_embeds, deepstack_visual_embeds = image_features
129
+ else:
130
+ image_embeds = image_features
131
+ image_embeds = torch.cat(image_embeds, dim=0).to(
132
+ inputs_embeds.device, inputs_embeds.dtype
133
+ ) # [N_v_total, d]
134
+ image_mask, _ = vl_model.get_placeholder_mask(
135
+ input_ids, inputs_embeds=inputs_embeds, image_features=image_embeds
136
+ )
137
+ inputs_embeds = inputs_embeds.masked_scatter(image_mask, image_embeds)
138
+ if is_qwen3_vl:
139
+ visual_pos_masks = image_mask[..., 0]
140
+
141
+ if is_qwen3_vl:
142
+ position_ids, _ = vl_model.get_rope_index(
143
+ input_ids,
144
+ image_grid_thw,
145
+ None, # video_grid_thw
146
+ attention_mask=attention_mask,
147
+ )
148
+ else:
149
+ position_ids, _ = vl_model.get_rope_index(
150
+ input_ids,
151
+ image_grid_thw,
152
+ None, # video_grid_thw
153
+ second_per_grid_ts=None,
154
+ attention_mask=attention_mask,
155
+ )
156
+ if return_deepstack:
157
+ return (
158
+ inputs_embeds,
159
+ position_ids,
160
+ visual_pos_masks,
161
+ deepstack_visual_embeds,
162
+ )
163
+ return inputs_embeds, position_ids
164
+
165
+
166
+ def make_split_context(
167
+ text_model,
168
+ inputs_embeds: torch.Tensor, # [B, L, d]
169
+ position_ids: torch.Tensor, # [3, B, L]
170
+ attention_mask: torch.Tensor | None = None,
171
+ past_key_values: Cache | None = None,
172
+ cache_position: torch.Tensor | None = None,
173
+ visual_pos_masks: torch.Tensor | None = None,
174
+ deepstack_visual_embeds: list[torch.Tensor] | None = None,
175
+ ) -> SplitContext:
176
+ """Build masks and rotary embeddings once, as the stock forward does.
177
+
178
+ Args:
179
+ text_model: ``Qwen2_5_VLTextModel`` (``vl_model.language_model``).
180
+ """
181
+ if cache_position is None:
182
+ past_seen = past_key_values.get_seq_length() if past_key_values is not None else 0
183
+ cache_position = torch.arange(
184
+ past_seen, past_seen + inputs_embeds.shape[1], device=inputs_embeds.device
185
+ ) # [L]
186
+
187
+ if position_ids.ndim == 2:
188
+ position_ids = position_ids[None, ...].expand(3, position_ids.shape[0], -1)
189
+
190
+ model_type = str(getattr(text_model.config, "model_type", ""))
191
+
192
+ # Packed-sequence convention: a leading text-only row makes it [4, B, L].
193
+ if position_ids.ndim == 3 and position_ids.shape[0] == 4:
194
+ text_position_ids = position_ids[0] # [B, L]
195
+ position_ids = position_ids[1:] # [3, B, L]
196
+ elif model_type == "qwen3_vl_text":
197
+ # Qwen3-VL always passes the temporal M-RoPE row to both the causal-mask
198
+ # builder and decoder layers, even for ordinary (non-packed) inputs.
199
+ text_position_ids = position_ids[0]
200
+ else:
201
+ text_position_ids = None
202
+
203
+ mask_kwargs: dict[str, Any] = {
204
+ "config": text_model.config,
205
+ "input_embeds": inputs_embeds,
206
+ "attention_mask": attention_mask,
207
+ "cache_position": cache_position,
208
+ "past_key_values": past_key_values,
209
+ "position_ids": text_position_ids,
210
+ }
211
+ causal_mask_mapping = {"full_attention": create_causal_mask(**mask_kwargs)}
212
+ if getattr(text_model, "has_sliding_layers", False):
213
+ causal_mask_mapping["sliding_attention"] = create_sliding_window_causal_mask(
214
+ **mask_kwargs
215
+ )
216
+
217
+ position_embeddings = text_model.rotary_emb(inputs_embeds, position_ids)
218
+
219
+ return SplitContext(
220
+ hidden_states=inputs_embeds,
221
+ position_ids=position_ids,
222
+ position_embeddings=position_embeddings,
223
+ causal_mask_mapping=causal_mask_mapping,
224
+ cache_position=cache_position,
225
+ text_position_ids=text_position_ids,
226
+ past_key_values=past_key_values,
227
+ attention_mask=attention_mask,
228
+ visual_pos_masks=visual_pos_masks,
229
+ deepstack_visual_embeds=deepstack_visual_embeds,
230
+ )
231
+
232
+
233
+ def block_attention_edges(
234
+ ctx: SplitContext,
235
+ query_mask: torch.Tensor,
236
+ key_mask: torch.Tensor,
237
+ ) -> SplitContext:
238
+ """Return ``ctx`` with selected query-to-key attention edges removed.
239
+
240
+ ``query_mask`` and ``key_mask`` are boolean ``[B, L]`` supports in the
241
+ current (cache-free) sequence. Every ordinary causal/padding constraint is
242
+ preserved; only their Cartesian product is additionally masked. Both
243
+ boolean SDPA masks (``True`` means visible) and additive eager masks
244
+ (``0``/negative infinity) are supported.
245
+
246
+ The helper deliberately rejects cached contexts. Its intended use is a
247
+ counterfactual recurrent layer evaluation, never autoregressive decoding,
248
+ and silently guessing the key offset of a populated cache would invalidate
249
+ the causal comparison.
250
+ """
251
+ if ctx.past_key_values is not None and ctx.past_key_values.get_seq_length() > 0:
252
+ raise ValueError("block_attention_edges requires a cache-free context")
253
+ if query_mask.dtype != torch.bool or key_mask.dtype != torch.bool:
254
+ raise TypeError("query_mask and key_mask must be boolean tensors")
255
+ if query_mask.shape != key_mask.shape or query_mask.ndim != 2:
256
+ raise ValueError(
257
+ "query_mask and key_mask must have the same [B, L] shape, got "
258
+ f"{tuple(query_mask.shape)} and {tuple(key_mask.shape)}"
259
+ )
260
+ batch_size, seq_len = query_mask.shape
261
+ if ctx.hidden_states.shape[:2] != (batch_size, seq_len):
262
+ raise ValueError(
263
+ "edge masks must match the SplitContext sequence, got "
264
+ f"{tuple(query_mask.shape)} for {tuple(ctx.hidden_states.shape[:2])}"
265
+ )
266
+
267
+ blocked = query_mask[:, None, :, None] & key_mask[:, None, None, :]
268
+ updated: dict[str, torch.Tensor] = {}
269
+ for attention_type, base_mask in ctx.causal_mask_mapping.items():
270
+ if base_mask is None:
271
+ # SDPA may omit an all-valid causal mask. Materialize exactly that
272
+ # lower triangle, then reapply key padding before deleting edges.
273
+ q_positions = ctx.cache_position
274
+ if q_positions.numel() != seq_len:
275
+ raise ValueError(
276
+ "cache-free context must have one cache position per row"
277
+ )
278
+ key_positions = torch.arange(seq_len, device=query_mask.device)
279
+ visible = key_positions[None, :] <= q_positions[:, None]
280
+ visible = visible[None, None].expand(batch_size, 1, -1, -1)
281
+ if ctx.attention_mask is not None:
282
+ if ctx.attention_mask.shape != (batch_size, seq_len):
283
+ raise ValueError(
284
+ "counterfactual edge masking expects a 2-D [B, L] "
285
+ "attention mask"
286
+ )
287
+ visible = visible & ctx.attention_mask[:, None, None, :].bool()
288
+ updated[attention_type] = visible & ~blocked
289
+ continue
290
+
291
+ if not isinstance(base_mask, torch.Tensor) or base_mask.ndim != 4:
292
+ raise TypeError(
293
+ "counterfactual edge masking supports tensor 4-D attention "
294
+ f"masks, got {type(base_mask)!r}"
295
+ )
296
+ if base_mask.shape[0] not in (1, batch_size):
297
+ raise ValueError("attention-mask batch dimension is incompatible")
298
+ if base_mask.shape[-2:] != (seq_len, seq_len):
299
+ raise ValueError(
300
+ "counterfactual edge masking expects a square cache-free mask, "
301
+ f"got {tuple(base_mask.shape)}"
302
+ )
303
+ if base_mask.dtype == torch.bool:
304
+ updated[attention_type] = base_mask & ~blocked
305
+ elif base_mask.is_floating_point():
306
+ updated[attention_type] = base_mask.masked_fill(
307
+ blocked, torch.finfo(base_mask.dtype).min
308
+ )
309
+ else:
310
+ raise TypeError(
311
+ f"unsupported attention mask dtype {base_mask.dtype}"
312
+ )
313
+
314
+ return replace(ctx, causal_mask_mapping=updated)
315
+
316
+
317
+ # ---------------------------------------------------------------------------
318
+ # running layer ranges
319
+ # ---------------------------------------------------------------------------
320
+
321
+
322
+ def run_layer_range(
323
+ text_model,
324
+ ctx: SplitContext,
325
+ start: int,
326
+ stop: int | None = None,
327
+ use_cache: bool = False,
328
+ hidden_states: torch.Tensor | None = None,
329
+ collect: bool = False,
330
+ ) -> torch.Tensor | tuple[torch.Tensor, list[torch.Tensor]]:
331
+ """Run ``text_model.layers[start:stop]`` on ``ctx``.
332
+
333
+ ``self.norm`` is *not* applied -- it belongs to the very top of the stack.
334
+ Call :func:`final_norm` after the last range.
335
+
336
+ Args:
337
+ hidden_states: override the context's states (leave ``None`` to chain).
338
+ collect: also return the input hidden states of every layer in the range
339
+ plus the range output, i.e. ``stop - start + 1`` tensors.
340
+
341
+ Returns:
342
+ ``[B, L, d]``, or ``(output, collected)`` when ``collect``.
343
+ """
344
+ layers = text_model.layers
345
+ stop = len(layers) if stop is None else stop
346
+ h = ctx.hidden_states if hidden_states is None else hidden_states
347
+
348
+ collected: list[torch.Tensor] = []
349
+ for layer_index, layer in enumerate(layers[start:stop], start=start):
350
+ if collect:
351
+ collected.append(h)
352
+ attention_type = getattr(layer, "attention_type", "full_attention")
353
+ h = layer(
354
+ h,
355
+ attention_mask=ctx.causal_mask_mapping[attention_type],
356
+ position_ids=ctx.text_position_ids,
357
+ past_key_values=ctx.past_key_values,
358
+ use_cache=use_cache,
359
+ cache_position=ctx.cache_position,
360
+ position_embeddings=ctx.position_embeddings,
361
+ )
362
+ # 4.57 decoder layers return a bare tensor; older ones returned a tuple.
363
+ if isinstance(h, tuple):
364
+ h = h[0]
365
+ if (
366
+ ctx.deepstack_visual_embeds is not None
367
+ and layer_index < len(ctx.deepstack_visual_embeds)
368
+ ):
369
+ if ctx.visual_pos_masks is None:
370
+ raise ValueError("DeepStack features require visual_pos_masks")
371
+ h = text_model._deepstack_process(
372
+ h,
373
+ ctx.visual_pos_masks,
374
+ ctx.deepstack_visual_embeds[layer_index],
375
+ )
376
+
377
+ if collect:
378
+ collected.append(h)
379
+ return h, collected
380
+ return h
381
+
382
+
383
+ def final_norm(text_model, hidden_states: torch.Tensor) -> torch.Tensor:
384
+ """Apply the stack's final RMSNorm. ``[B, L, d] -> [B, L, d]``."""
385
+ return text_model.norm(hidden_states)
386
+
387
+
388
+ # ---------------------------------------------------------------------------
389
+ # token selection operators (Pi_img / Pi_q in §3.2)
390
+ # ---------------------------------------------------------------------------
391
+
392
+
393
+ def image_token_mask(input_ids: torch.LongTensor, image_token_id: int) -> torch.Tensor:
394
+ """``Pi_img`` support: ``[B, L]`` bool, True at image placeholder positions."""
395
+ return input_ids == image_token_id
396
+
397
+
398
+ def vision_span_mask(input_ids: torch.LongTensor, config) -> torch.Tensor:
399
+ """``[B, L]`` bool covering ``<|vision_start|>``, image pads, ``<|vision_end|>``.
400
+
401
+ Use this (not :func:`image_token_mask`) when *removing* the visual segment to
402
+ build the text-only branch, so the delimiters do not survive as orphans.
403
+ """
404
+ ids = {
405
+ config.vision_start_token_id,
406
+ config.vision_end_token_id,
407
+ config.image_token_id,
408
+ config.video_token_id,
409
+ }
410
+ mask = torch.zeros_like(input_ids, dtype=torch.bool)
411
+ for tid in ids:
412
+ mask |= input_ids == tid
413
+ return mask
414
+
415
+
416
+ def select_tokens(
417
+ hidden_states: torch.Tensor, # [B, L, d]
418
+ mask: torch.Tensor, # [B, L] bool
419
+ ) -> torch.Tensor:
420
+ """Gather masked positions. Requires an equal count per batch element.
421
+
422
+ Returns ``[B, N, d]`` where ``N`` is that per-element count.
423
+ """
424
+ counts = mask.sum(dim=1)
425
+ if counts.numel() > 1 and not bool((counts == counts[0]).all()):
426
+ raise ValueError(
427
+ f"select_tokens needs the same number of selected tokens per batch "
428
+ f"element, got {counts.tolist()}. Bucket by visual-token count or "
429
+ f"gather per-example instead."
430
+ )
431
+ n = int(counts[0])
432
+ b, _, d = hidden_states.shape
433
+ return hidden_states[mask].view(b, n, d)
434
+
435
+
436
+ def select_tokens_padded(
437
+ hidden_states: torch.Tensor, # [B, L, d]
438
+ mask: torch.Tensor, # [B, L] bool
439
+ ) -> tuple[torch.Tensor, torch.Tensor]:
440
+ """Gather masked positions, right-padded to the batch maximum.
441
+
442
+ Qwen2.5-VL uses dynamic resolution, so ``N_v`` differs across a batch. §3.1's
443
+ patching genuinely needs equal counts (it transplants position by position),
444
+ but ``r_theta`` only cross-attends over ``V*`` -- a variable-length memory is
445
+ exactly what a key-padding mask is for.
446
+
447
+ Returns ``(padded [B, N_max, d], key_padding_mask [B, N_max])`` where the
448
+ mask is ``True`` at padding, matching ``nn.MultiheadAttention``.
449
+ """
450
+ counts = mask.sum(dim=1)
451
+ n_max = int(counts.max())
452
+ b, _, d = hidden_states.shape
453
+ out = hidden_states.new_zeros((b, n_max, d))
454
+ pad = torch.ones((b, n_max), dtype=torch.bool, device=hidden_states.device)
455
+ for i in range(b):
456
+ n = int(counts[i])
457
+ out[i, :n] = hidden_states[i][mask[i]]
458
+ pad[i, :n] = False
459
+ return out, pad
460
+
461
+
462
+ __all__ = [
463
+ "SplitContext",
464
+ "embed_multimodal",
465
+ "make_split_context",
466
+ "block_attention_edges",
467
+ "run_layer_range",
468
+ "final_norm",
469
+ "image_token_mask",
470
+ "vision_span_mask",
471
+ "select_tokens",
472
+ "select_tokens_padded",
473
+ ]