Visual Question Answering
Transformers
Safetensors
cvrr_merged
feature-extraction
cvrr
custom_code
latent-reasoning
Instructions to use dmis-lab/InternVL3-9B-CVRR with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use dmis-lab/InternVL3-9B-CVRR with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("visual-question-answering", model="dmis-lab/InternVL3-9B-CVRR", trust_remote_code=True)# Load model directly from transformers import AutoModel model = AutoModel.from_pretrained("dmis-lab/InternVL3-9B-CVRR", trust_remote_code=True, device_map="auto") - Notebooks
- Google Colab
- Kaggle
Add files using upload-large-folder tool
Browse files- NOTICE.txt +4 -0
- README.md +121 -0
- config.json +30 -0
- configuration_cvrr_merged.py +18 -0
- configuration_source_qwen25.py +1142 -0
- configuration_source_qwen3.py +740 -0
- cvrr_release_config.json +19 -0
- licenses/Apache-2.0.txt +202 -0
- licenses/InternVL-MIT.txt +21 -0
- licenses/UPSTREAM_MODEL_CARD.md +701 -0
- merged_transition.safetensors +3 -0
- modeling_cvrr_merged.py +279 -0
- modeling_source_qwen25.py +0 -0
- modeling_source_qwen3.py +0 -0
- native_backbone/added_tokens.json +11 -0
- native_backbone/config.json +225 -0
- native_backbone/configuration_intern_vit.py +119 -0
- native_backbone/configuration_internlm2.py +150 -0
- native_backbone/configuration_internvl_chat.py +101 -0
- native_backbone/conversation.py +391 -0
- native_backbone/generation_config.json +4 -0
- native_backbone/model.safetensors.index.json +692 -0
- native_backbone/modeling_intern_vit.py +429 -0
- native_backbone/modeling_internlm2.py +1456 -0
- native_backbone/modeling_internvl_chat.py +363 -0
- native_backbone/native-00001.safetensors +3 -0
- native_backbone/native-00002.safetensors +3 -0
- native_backbone/native-00003.safetensors +3 -0
- native_backbone/native-00004.safetensors +3 -0
- native_backbone/special_tokens_map.json +63 -0
- native_backbone/tokenization_internlm3.py +294 -0
- native_backbone/tokenizer.model +3 -0
- native_backbone/tokenizer_config.json +330 -0
- requirements.txt +11 -0
- source_gemma.py +685 -0
- source_helpers.py +69 -0
- source_internvl.py +286 -0
- source_perceive.py +414 -0
- source_spatial.py +349 -0
- 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 |
+

|
| 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 |
+

|
| 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 |
+

|
| 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 |
+

|
| 107 |
+
|
| 108 |
+
### OCR, Chart, and Document Understanding
|
| 109 |
+
|
| 110 |
+

|
| 111 |
+
|
| 112 |
+
### Multi-Image & Real-World Comprehension
|
| 113 |
+
|
| 114 |
+

|
| 115 |
+
|
| 116 |
+
### Comprehensive Multimodal & Hallucination Evaluation
|
| 117 |
+
|
| 118 |
+

|
| 119 |
+
|
| 120 |
+
### Visual Grounding
|
| 121 |
+
|
| 122 |
+

|
| 123 |
+
|
| 124 |
+
### Multimodal Multilingual Understanding
|
| 125 |
+
|
| 126 |
+

|
| 127 |
+
|
| 128 |
+
### Video Understanding
|
| 129 |
+
|
| 130 |
+

|
| 131 |
+
|
| 132 |
+
### GUI Grounding
|
| 133 |
+
|
| 134 |
+

|
| 135 |
+
|
| 136 |
+
### Spatial Reasoning
|
| 137 |
+
|
| 138 |
+

|
| 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 |
+

|
| 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 |
+

|
| 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 |
+

|
| 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 |
+

|
| 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 |
+
]
|