manu02 commited on
Commit
f896f0e
·
verified ·
1 Parent(s): 6ea3099

Upload benchmarked LANA model

Browse files
.gitattributes CHANGED
@@ -33,3 +33,4 @@ saved_model/**/* filter=lfs diff=lfs merge=lfs -text
33
  *.zip filter=lfs diff=lfs merge=lfs -text
34
  *.zst filter=lfs diff=lfs merge=lfs -text
35
  *tfevents* filter=lfs diff=lfs merge=lfs -text
 
 
33
  *.zip filter=lfs diff=lfs merge=lfs -text
34
  *.zst filter=lfs diff=lfs merge=lfs -text
35
  *tfevents* filter=lfs diff=lfs merge=lfs -text
36
+ assets/AnatomicalAttention.gif filter=lfs diff=lfs merge=lfs -text
DINOv3-LICENSE.txt ADDED
@@ -0,0 +1,65 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ DINOv3 License
2
+
3
+ Last Updated: August 19, 2025
4
+
5
+ "Agreement" means the terms and conditions for use, reproduction, distribution and modification of the DINO Materials set forth herein.
6
+
7
+ "DINO Materials" means, collectively, Documentation and the models, software and algorithms, including machine-learning model code, trained model weights, inference-enabling code, training-enabling code, fine-tuning enabling code, and other elements of the foregoing distributed by Meta and made available under this Agreement.
8
+
9
+ "Documentation" means the specifications, manuals and documentation accompanying DINO Materials distributed by Meta.
10
+
11
+ "Licensee" or "you" means you, or your employer or any other person or entity (if you are entering into this Agreement on such person or entity’s behalf), of the age required under applicable laws, rules or regulations to provide legal consent and that has legal authority to bind your employer or such other person or entity if you are entering in this Agreement on their behalf.
12
+
13
+ "Meta" or "we" means Meta Platforms Ireland Limited (if you are located in or, if you are an entity, your principal place of business is in the EEA or Switzerland) or Meta Platforms, Inc. (if you are located outside of the EEA or Switzerland).
14
+
15
+ "Sanctions" means any economic or trade sanctions or restrictions administered or enforced by the United States (including the Office of Foreign Assets Control of the U.S. Department of the Treasury ("OFAC"), the U.S. Department of State and the U.S. Department of Commerce), the United Nations, the European Union, or the United Kingdom.
16
+
17
+ "Trade Controls" means any of the following: Sanctions and applicable export and import controls.
18
+
19
+ By clicking "I Accept" below or by using or distributing any portion or element of the DINO Materials, you agree to be bound by this Agreement.
20
+
21
+ 1. License Rights and Redistribution.
22
+
23
+ a. Grant of Rights. You are granted a non-exclusive, worldwide, non-transferable and royalty-free limited license under Meta's intellectual property or other rights owned by Meta embodied in the DINO Materials to use, reproduce, distribute, copy, create derivative works of, and make modifications to the DINO Materials.
24
+
25
+ b. Redistribution and Use.
26
+ i. Distribution of DINO Materials, and any derivative works thereof, are subject to the terms of this Agreement. If you distribute or make the DINO Materials, or any derivative works thereof, available to a third party, you may only do so under the terms of this Agreement and you shall provide a copy of this Agreement with any such DINO Materials.
27
+ ii. If you submit for publication the results of research you perform on, using, or otherwise in connection with DINO Materials, you must acknowledge the use of DINO Materials in your publication.
28
+ iii. Your use of the DINO Materials must comply with applicable laws and regulations, including Trade Control Laws and applicable privacy and data protection laws.
29
+ iv. Your use of the DINO Materials will not involve or encourage others to reverse engineer, decompile or discover the underlying components of the DINO Materials.
30
+ v. You are not the target of Trade Controls and your use of DINO Materials must comply with Trade Controls. You agree not to use, or permit others to use, DINO Materials for any activities subject to the International Traffic in Arms Regulations (ITAR) or end uses prohibited by Trade Controls, including those related to military or warfare purposes, nuclear industries or applications, espionage, or the development or use of guns or illegal weapons.
31
+
32
+ 2. User Support.
33
+
34
+ Your use of the DINO Materials is done at your own discretion; Meta does not process any information nor provide any service in relation to such use. Meta is under no obligation to provide any support services for the DINO Materials. Any support provided is "as is", "with all faults", and without warranty of any kind.
35
+
36
+ 3. Disclaimer of Warranty.
37
+
38
+ UNLESS REQUIRED BY APPLICABLE LAW, THE DINO MATERIALS AND ANY OUTPUT AND RESULTS THEREFROM ARE PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, AND META DISCLAIMS ALL WARRANTIES OF ANY KIND, BOTH EXPRESS AND IMPLIED, INCLUDING, WITHOUT LIMITATION, ANY WARRANTIES OF TITLE, NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
39
+
40
+ YOU ARE SOLELY RESPONSIBLE FOR DETERMINING THE APPROPRIATENESS OF USING OR REDISTRIBUTING THE DINO MATERIALS AND ASSUME ANY RISKS ASSOCIATED WITH YOUR USE OF THE DINO MATERIALS AND ANY OUTPUT AND RESULTS.
41
+
42
+ 4. Limitation of Liability.
43
+
44
+ IN NO EVENT WILL META OR ITS AFFILIATES BE LIABLE UNDER ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, TORT, NEGLIGENCE, PRODUCTS LIABILITY, OR OTHERWISE, ARISING OUT OF THIS AGREEMENT, FOR ANY LOST PROFITS OR ANY DIRECT OR INDIRECT, SPECIAL, CONSEQUENTIAL, INCIDENTAL, EXEMPLARY OR PUNITIVE DAMAGES, EVEN IF META OR ITS AFFILIATES HAVE BEEN ADVISED OF THE POSSIBILITY OF ANY OF THE FOREGOING.
45
+
46
+ 5. Intellectual Property.
47
+
48
+ a. Subject to Meta's ownership of DINO Materials and derivatives made by or for Meta, with respect to any derivative works and modifications of the DINO Materials that are made by you, as between you and Meta, you are and will be the owner of such derivative works and modifications.
49
+ b. If you institute litigation or other proceedings against Meta or any entity (including a cross-claim or counterclaim in a lawsuit) alleging that the DINO Materials, outputs or results, or any portion of any of the foregoing, constitutes infringement of intellectual property or other rights owned or licensable by you, then any licenses granted to you under this Agreement shall terminate as of the date such litigation or claim is filed or instituted.
50
+
51
+ You will indemnify and hold harmless Meta from and against any claim by any third party arising out of or related to your use or distribution of the DINO Materials.
52
+
53
+ 6. Term and Termination.
54
+
55
+ The term of this Agreement will commence upon your acceptance of this Agreement or access to the DINO Materials and will continue in full force and effect until terminated in accordance with the terms and conditions herein. Meta may terminate this Agreement if you are in breach of any term or condition of this Agreement. Upon termination of this Agreement, you shall delete and cease use of the DINO Materials. Sections 3, 4 and 7 shall survive the termination of this Agreement.
56
+
57
+ 7. Governing Law and Jurisdiction.
58
+
59
+ This Agreement will be governed and construed under the laws of the State of California without regard to choice of law principles, and the UN Convention on Contracts for the International Sale of Goods does not apply to this Agreement. The courts of California shall have exclusive jurisdiction of any dispute arising out of this Agreement.
60
+
61
+ 8. Modifications and Amendments.
62
+
63
+ Meta may modify this Agreement from time to time; provided that they are similar in spirit to the current version of the Agreement, but may differ in detail to address new problems or concerns. All such changes will be effective immediately. Your continued use of the DINO Materials after any modification to this Agreement constitutes your agreement to such modification.
64
+
65
+ Except as provided in this Agreement, no modification or addition to any provision of this Agreement will be binding unless it is in writing and signed by an authorized representative of both you and Meta.
README.md ADDED
@@ -0,0 +1,186 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ license: mit
3
+ library_name: transformers
4
+ pipeline_tag: image-to-text
5
+ tags:
6
+ - medical-ai
7
+ - radiology
8
+ - chest-xray
9
+ - report-generation
10
+ - segmentation
11
+ - anatomical-attention
12
+ metrics:
13
+ - BLEU
14
+ - METEOR
15
+ - ROUGE
16
+ - CIDEr
17
+ ---
18
+
19
+ # LAnA
20
+
21
+ **Layer-Wise Anatomical Attention model**
22
+
23
+ > Best current model in this collection: [`manu02/LAnA-v3`](https://huggingface.co/manu02/LAnA-v3)
24
+
25
+ [![ArXiv](https://img.shields.io/badge/ArXiv-2512.16841-B31B1B?logo=arxiv&logoColor=white)](https://arxiv.org/abs/2512.16841)
26
+ [![LinkedIn](https://img.shields.io/badge/LinkedIn-devmuniz-0A66C2?logo=linkedin&logoColor=white)](https://www.linkedin.com/in/devmuniz)
27
+ [![GitHub Profile](https://img.shields.io/badge/GitHub-devMuniz02-181717?logo=github&logoColor=white)](https://github.com/devMuniz02)
28
+ [![Portfolio](https://img.shields.io/badge/Portfolio-devmuniz02.github.io-0F172A?logo=googlechrome&logoColor=white)](https://devmuniz02.github.io/)
29
+ [![GitHub Repo](https://img.shields.io/badge/Repository-layer--wise--anatomical--attention-181717?logo=github&logoColor=white)](https://github.com/devMuniz02/layer-wise-anatomical-attention)
30
+ [![Hugging Face](https://img.shields.io/badge/Hugging%20Face-manu02-FFD21E?logoColor=black)](https://huggingface.co/manu02)
31
+
32
+ ![Layer-Wise Anatomical Attention](assets/AnatomicalAttention.gif)
33
+
34
+ ## Overview
35
+
36
+ LAnA is a medical report-generation project for chest X-ray images. The completed project is intended to generate radiology reports with a vision-language model guided by layer-wise anatomical attention built from predicted anatomical masks.
37
+
38
+ The architecture combines a DINOv3 vision encoder, lung and heart segmentation heads, and a GPT-2 decoder modified so each transformer layer receives a different anatomical attention bias derived from the segmentation mask.
39
+
40
+ ## Intended Use
41
+
42
+ - Input: a chest X-ray image resized to `512x512` and normalized with ImageNet mean/std.
43
+ - Output: a generated radiology report.
44
+ - Best fit: research use, report-generation experiments, and anatomical-attention ablations.
45
+
46
+ ## How to Run
47
+
48
+ New users should prefer the standard Hugging Face flow below.
49
+ The legacy snapshot/manual implementation lives on the `snapshot-legacy` branch for backward compatibility.
50
+
51
+ ### Implementation 1: Standard Hugging Face loading
52
+
53
+ ```python
54
+ import torch
55
+ from PIL import Image
56
+ from transformers import AutoModel, AutoProcessor
57
+
58
+ repo_id = "manu02/LAnA-v5"
59
+ device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
60
+
61
+ processor = AutoProcessor.from_pretrained(repo_id, trust_remote_code=True)
62
+ model = AutoModel.from_pretrained(repo_id, trust_remote_code=True)
63
+ model.move_non_quantized_modules(device)
64
+ model.eval()
65
+
66
+ image = Image.open("example.png").convert("RGB")
67
+ inputs = processor(images=image, return_tensors="pt")
68
+ inputs = {name: tensor.to(device) for name, tensor in inputs.items()}
69
+
70
+ with torch.inference_mode():
71
+ generated = model.generate(**inputs, max_new_tokens=150)
72
+
73
+ report = processor.batch_decode(generated, skip_special_tokens=True)[0]
74
+ print(report)
75
+ ```
76
+
77
+ Batched inference uses the same path:
78
+
79
+ ```python
80
+ batch = processor(images=[image_a, image_b], return_tensors="pt")
81
+ batch = {name: tensor.to(device) for name, tensor in batch.items()}
82
+ generated = model.generate(**batch, max_new_tokens=150)
83
+ reports = processor.batch_decode(generated, skip_special_tokens=True)
84
+ ```
85
+
86
+ `HF_TOKEN` is optional for this public standard-loading path. If you do not set one, the model still loads,
87
+ but Hugging Face may show lower-rate-limit warnings.
88
+
89
+ ### Legacy snapshot branch
90
+
91
+ Use the snapshot/manual branch only if you specifically need the older import-based workflow:
92
+
93
+ - Branch: [`snapshot-legacy`](https://huggingface.co/manu02/LAnA-v5/tree/snapshot-legacy)
94
+ - Download example: `snapshot_download("manu02/LAnA-v5", revision="snapshot-legacy")`
95
+
96
+ ## Licensing and Redistribution Notice
97
+
98
+ This checkpoint bundles or derives from Meta DINOv3 model materials. Redistribution of those components must follow
99
+ the DINOv3 license terms included in this repository. The project code remains available under the repository's own
100
+ license, but the full packaged checkpoint should not be treated as MIT-only.
101
+
102
+ ## Research and Safety Disclaimer
103
+
104
+ This model is intended for research and educational use only. It is not a medical device, has not been validated
105
+ for clinical deployment, and should not be used as a substitute for professional radiology review.
106
+
107
+ ## MIMIC Test Results
108
+
109
+ Frontal-only evaluation using `PA/AP` studies only.
110
+
111
+ ### Current Checkpoint Results
112
+
113
+ | Metric | Value |
114
+ | --- | --- |
115
+ | Number of studies | TBD |
116
+ | RadGraph F1 | TBD |
117
+ | RadGraph entity F1 | TBD |
118
+ | RadGraph relation F1 | TBD |
119
+ | CheXpert F1 14-micro | TBD |
120
+ | CheXpert F1 5-micro | TBD |
121
+ | CheXpert F1 14-macro | TBD |
122
+ | CheXpert F1 5-macro | TBD |
123
+
124
+ ### Final Completed Training Results
125
+
126
+ The final table will be populated when the planned training run is completed. Until then, final-report metrics remain `TBD`.
127
+
128
+ | Metric | Value |
129
+ | --- | --- |
130
+ | Number of studies | TBD |
131
+ | RadGraph F1 | TBD |
132
+ | RadGraph entity F1 | TBD |
133
+ | RadGraph relation F1 | TBD |
134
+ | CheXpert F1 14-micro | TBD |
135
+ | CheXpert F1 5-micro | TBD |
136
+ | CheXpert F1 14-macro | TBD |
137
+ | CheXpert F1 5-macro | TBD |
138
+
139
+
140
+ ## Data
141
+
142
+ - Full project datasets: CheXpert and MIMIC-CXR.
143
+ - Intended project scope: train on curated chest X-ray/report data from both datasets and evaluate on MIMIC-CXR test studies.
144
+ - Current released checkpoint datasets: `MIMIC-CXR (findings-only)` for training and `MIMIC-CXR (findings-only)` for validation.
145
+ - Current published evaluation: MIMIC-CXR test split, `frontal-only (PA/AP)` studies.
146
+
147
+ ## Evaluation
148
+
149
+ - Medical report metrics implemented in the repository include RadGraph F1 and CheXpert F1 (`14-micro`, `5-micro`, `14-macro`, `5-macro`).
150
+
151
+ ## Training Snapshot
152
+
153
+ - Run: `LAnA-v5`
154
+ - This section describes the current public checkpoint, not the final completed project.
155
+ - Method: `full_adamw`
156
+ - Vision encoder: `facebook/dinov3-vits16-pretrain-lvd1689m`
157
+ - Text decoder: `gpt2`
158
+ - Visual projection: `linear`
159
+ - Segmentation encoder: `facebook/dinov3-convnext-small-pretrain-lvd1689m`
160
+ - Image size: `512`
161
+ - Local batch size: `1`
162
+ - Effective global batch size: `16`
163
+ - Scheduler: `cosine`
164
+ - Warmup steps: `1318`
165
+ - Weight decay: `0.01`
166
+ - Steps completed: `482`
167
+ - Planned total steps: `26358`
168
+ - Images seen: `7724`
169
+ - Total training time: `0.1667` hours
170
+ - Hardware: `NVIDIA GeForce RTX 5070`
171
+ - Final train loss: `8.1020`
172
+ - Validation loss: `8.1601`
173
+
174
+ ## Status
175
+
176
+ - Project status: `Training in progress`
177
+ - Release status: `Research preview checkpoint`
178
+ - Current checkpoint status: `Not final`
179
+ - Training completion toward planned run: `1.83%` (`0` / `3` epochs)
180
+ - Current published metrics are intermediate and will change as training continues.
181
+
182
+ ## Notes
183
+
184
+ - Set `HF_TOKEN` with permission to access the DINOv3 repositories required by this model before downloading or running inference.
185
+ - `segmenters/` contains the lung and heart segmentation checkpoints used to build anatomical attention masks.
186
+ - `evaluations/mimic_test_metrics.json` contains the latest saved MIMIC test metrics.
__init__.py ADDED
@@ -0,0 +1,13 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from .configuration_lana import LanaConfig
2
+ from .image_processing_lana import LanaImageProcessor
3
+ from .modeling_lana import LanaForConditionalGeneration
4
+ from .modeling_outputs import LanaModelOutput
5
+ from .processing_lana import LanaProcessor
6
+
7
+ __all__ = [
8
+ "LanaConfig",
9
+ "LanaImageProcessor",
10
+ "LanaForConditionalGeneration",
11
+ "LanaModelOutput",
12
+ "LanaProcessor",
13
+ ]
assets/AnatomicalAttention.gif ADDED

Git LFS Details

  • SHA256: 3854885a631419336dca34b3375e29be91597a994349e3abdf1460e6908ec391
  • Pointer size: 133 Bytes
  • Size of remote file: 27.3 MB
bundled_backbones/segmenter_encoder/config.json ADDED
@@ -0,0 +1,27 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "architectures": [
3
+ "DINOv3ConvNextModel"
4
+ ],
5
+ "depths": [
6
+ 3,
7
+ 3,
8
+ 27,
9
+ 3
10
+ ],
11
+ "drop_path_rate": 0.0,
12
+ "hidden_act": "gelu",
13
+ "hidden_sizes": [
14
+ 96,
15
+ 192,
16
+ 384,
17
+ 768
18
+ ],
19
+ "image_size": 224,
20
+ "initializer_range": 0.02,
21
+ "layer_norm_eps": 1e-06,
22
+ "layer_scale_init_value": 1e-06,
23
+ "model_type": "dinov3_convnext",
24
+ "num_channels": 3,
25
+ "torch_dtype": "float32",
26
+ "transformers_version": "4.56.0.dev0"
27
+ }
bundled_backbones/text_decoder/config.json ADDED
@@ -0,0 +1,31 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "activation_function": "gelu_new",
3
+ "architectures": [
4
+ "GPT2LMHeadModel"
5
+ ],
6
+ "attn_pdrop": 0.1,
7
+ "bos_token_id": 50256,
8
+ "embd_pdrop": 0.1,
9
+ "eos_token_id": 50256,
10
+ "initializer_range": 0.02,
11
+ "layer_norm_epsilon": 1e-05,
12
+ "model_type": "gpt2",
13
+ "n_ctx": 1024,
14
+ "n_embd": 768,
15
+ "n_head": 12,
16
+ "n_layer": 12,
17
+ "n_positions": 1024,
18
+ "resid_pdrop": 0.1,
19
+ "summary_activation": null,
20
+ "summary_first_dropout": 0.1,
21
+ "summary_proj_to_labels": true,
22
+ "summary_type": "cls_index",
23
+ "summary_use_proj": true,
24
+ "task_specific_params": {
25
+ "text-generation": {
26
+ "do_sample": true,
27
+ "max_length": 50
28
+ }
29
+ },
30
+ "vocab_size": 50257
31
+ }
bundled_backbones/vision_encoder/config.json ADDED
@@ -0,0 +1,32 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "architectures": [
3
+ "DINOv3ViTModel"
4
+ ],
5
+ "attention_dropout": 0.0,
6
+ "drop_path_rate": 0.0,
7
+ "hidden_act": "gelu",
8
+ "hidden_size": 384,
9
+ "image_size": 224,
10
+ "initializer_range": 0.02,
11
+ "intermediate_size": 1536,
12
+ "key_bias": false,
13
+ "layer_norm_eps": 1e-05,
14
+ "layerscale_value": 1.0,
15
+ "mlp_bias": true,
16
+ "model_type": "dinov3_vit",
17
+ "num_attention_heads": 6,
18
+ "num_channels": 3,
19
+ "num_hidden_layers": 12,
20
+ "num_register_tokens": 4,
21
+ "patch_size": 16,
22
+ "pos_embed_jitter": null,
23
+ "pos_embed_rescale": 2.0,
24
+ "pos_embed_shift": null,
25
+ "proj_bias": true,
26
+ "query_bias": true,
27
+ "rope_theta": 100.0,
28
+ "torch_dtype": "float32",
29
+ "transformers_version": "4.56.0.dev0",
30
+ "use_gated_mlp": false,
31
+ "value_bias": true
32
+ }
config.json ADDED
@@ -0,0 +1,46 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "anatomical_attention_bias": 2.0,
3
+ "architectures": [
4
+ "LanaForConditionalGeneration"
5
+ ],
6
+ "attention_bias_mode": "gaussian_legacy",
7
+ "bundled_segmentation_model_name": "bundled_backbones/segmenter_encoder",
8
+ "bundled_text_model_name": "bundled_backbones/text_decoder",
9
+ "bundled_tokenizer_name": ".",
10
+ "bundled_vision_model_name": "bundled_backbones/vision_encoder",
11
+ "decoder_compute_dtype": "bfloat16",
12
+ "decoder_load_in_4bit": false,
13
+ "dtype": "float32",
14
+ "freeze_segmenter": true,
15
+ "generation_repetition_penalty": 1.2,
16
+ "generation_stop_on_eos": true,
17
+ "generation_use_bos_token": false,
18
+ "heart_segmenter_checkpoint": "segmenters/heart_segmenter_dinounet_best.pth",
19
+ "image_size": 512,
20
+ "layer_mask_base_kernel_size": 3,
21
+ "layer_mask_kernel_growth": 2,
22
+ "local_repo_path": "",
23
+ "lung_segmenter_checkpoint": "segmenters/lung_segmenter_dinounet_finetuned.pth",
24
+ "mask_size": 32,
25
+ "max_position_embeddings": 2048,
26
+ "model_type": "lana_radgen",
27
+ "num_attention_layers": 12,
28
+ "segmentation_attention_implementation": "sdpa",
29
+ "segmentation_model_name": "facebook/dinov3-convnext-small-pretrain-lvd1689m",
30
+ "segmenter_weights_in_model_state": true,
31
+ "text_hidden_size": 768,
32
+ "text_model_name": "gpt2",
33
+ "transformers_version": "5.3.0",
34
+ "use_cache": true,
35
+ "use_segmentation_mask": true,
36
+ "vision_feature_prefix_tokens_to_skip": 5,
37
+ "vision_model_name": "facebook/dinov3-vits16-pretrain-lvd1689m",
38
+ "visual_feature_dim": 384,
39
+ "visual_projection_type": "linear",
40
+ "vocab_size": 50257,
41
+ "auto_map": {
42
+ "AutoConfig": "configuration_lana.LanaConfig",
43
+ "AutoModel": "modeling_lana.LanaForConditionalGeneration",
44
+ "AutoProcessor": "processing_lana.LanaProcessor"
45
+ }
46
+ }
configuration_lana.py ADDED
@@ -0,0 +1,98 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from pathlib import Path
2
+
3
+ from huggingface_hub import snapshot_download
4
+ from transformers import PretrainedConfig
5
+
6
+
7
+ class LanaConfig(PretrainedConfig):
8
+ model_type = "lana_radgen"
9
+
10
+ @classmethod
11
+ def from_pretrained(cls, pretrained_model_name_or_path, *args, **kwargs):
12
+ loaded = super().from_pretrained(pretrained_model_name_or_path, *args, **kwargs)
13
+ if isinstance(loaded, tuple):
14
+ config, unused_kwargs = loaded
15
+ else:
16
+ config, unused_kwargs = loaded, None
17
+ repo_path = str(pretrained_model_name_or_path)
18
+ if not Path(repo_path).exists():
19
+ try:
20
+ repo_path = snapshot_download(repo_path)
21
+ except Exception:
22
+ repo_path = str(pretrained_model_name_or_path)
23
+ config.local_repo_path = repo_path
24
+ if unused_kwargs is not None:
25
+ return config, unused_kwargs
26
+ return config
27
+
28
+ def __init__(
29
+ self,
30
+ vision_model_name: str = "facebook/dinov3-vits16-pretrain-lvd1689m",
31
+ text_model_name: str = "gpt2",
32
+ image_size: int = 512,
33
+ mask_size: int = 32,
34
+ num_attention_layers: int = 12,
35
+ max_position_embeddings: int = 2048,
36
+ visual_feature_dim: int = 384,
37
+ text_hidden_size: int = 768,
38
+ visual_projection_type: str = "mlp4",
39
+ vocab_size: int = 50257,
40
+ layer_mask_base_kernel_size: int = 3,
41
+ layer_mask_kernel_growth: int = 2,
42
+ anatomical_attention_bias: float = 2.0,
43
+ attention_bias_mode: str = "layerwise",
44
+ vision_feature_prefix_tokens_to_skip: int = 1,
45
+ use_segmentation_mask: bool = True,
46
+ segmentation_model_name: str = "facebook/dinov3-convnext-small-pretrain-lvd1689m",
47
+ segmentation_attention_implementation: str = "sdpa",
48
+ freeze_segmenter: bool = True,
49
+ generation_use_bos_token: bool = True,
50
+ generation_stop_on_eos: bool = False,
51
+ generation_repetition_penalty: float = 1.0,
52
+ lung_segmenter_checkpoint: str = "",
53
+ heart_segmenter_checkpoint: str = "",
54
+ bundled_vision_model_name: str = "",
55
+ bundled_segmentation_model_name: str = "",
56
+ bundled_text_model_name: str = "",
57
+ bundled_tokenizer_name: str = "",
58
+ segmenter_weights_in_model_state: bool = False,
59
+ local_repo_path: str = "",
60
+ use_cache: bool = True,
61
+ decoder_load_in_4bit: bool = False,
62
+ decoder_compute_dtype: str = "float16",
63
+ **kwargs,
64
+ ):
65
+ self.vision_model_name = vision_model_name
66
+ self.text_model_name = text_model_name
67
+ self.image_size = image_size
68
+ self.mask_size = mask_size
69
+ self.num_attention_layers = num_attention_layers
70
+ self.max_position_embeddings = max_position_embeddings
71
+ self.visual_feature_dim = visual_feature_dim
72
+ self.text_hidden_size = text_hidden_size
73
+ self.visual_projection_type = visual_projection_type
74
+ self.vocab_size = vocab_size
75
+ self.layer_mask_base_kernel_size = layer_mask_base_kernel_size
76
+ self.layer_mask_kernel_growth = layer_mask_kernel_growth
77
+ self.anatomical_attention_bias = anatomical_attention_bias
78
+ self.attention_bias_mode = attention_bias_mode
79
+ self.vision_feature_prefix_tokens_to_skip = vision_feature_prefix_tokens_to_skip
80
+ self.use_segmentation_mask = use_segmentation_mask
81
+ self.segmentation_model_name = segmentation_model_name
82
+ self.segmentation_attention_implementation = segmentation_attention_implementation
83
+ self.freeze_segmenter = freeze_segmenter
84
+ self.generation_use_bos_token = generation_use_bos_token
85
+ self.generation_stop_on_eos = generation_stop_on_eos
86
+ self.generation_repetition_penalty = generation_repetition_penalty
87
+ self.lung_segmenter_checkpoint = lung_segmenter_checkpoint
88
+ self.heart_segmenter_checkpoint = heart_segmenter_checkpoint
89
+ self.bundled_vision_model_name = bundled_vision_model_name
90
+ self.bundled_segmentation_model_name = bundled_segmentation_model_name
91
+ self.bundled_text_model_name = bundled_text_model_name
92
+ self.bundled_tokenizer_name = bundled_tokenizer_name
93
+ self.segmenter_weights_in_model_state = segmenter_weights_in_model_state
94
+ self.local_repo_path = local_repo_path
95
+ self.use_cache = use_cache
96
+ self.decoder_load_in_4bit = decoder_load_in_4bit
97
+ self.decoder_compute_dtype = decoder_compute_dtype
98
+ super().__init__(**kwargs)
gpt2_modified.py ADDED
@@ -0,0 +1,399 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from typing import Optional, Union
2
+ import inspect
3
+
4
+ import torch
5
+ import torch.nn.functional as F
6
+ from torch import nn
7
+ from transformers import GPT2Config, GPT2LMHeadModel, GPT2Model
8
+ from transformers.cache_utils import Cache, DynamicCache, EncoderDecoderCache
9
+ from transformers.masking_utils import create_causal_mask
10
+ from transformers.modeling_attn_mask_utils import _prepare_4d_attention_mask_for_sdpa
11
+ from transformers.modeling_outputs import BaseModelOutputWithPastAndCrossAttentions, CausalLMOutputWithCrossAttentions
12
+ from transformers.modeling_utils import ALL_ATTENTION_FUNCTIONS
13
+ from transformers.models.gpt2.modeling_gpt2 import GPT2Attention, GPT2Block, eager_attention_forward
14
+
15
+ _CREATE_CAUSAL_MASK_EMBEDS_ARG = "inputs_embeds" if "inputs_embeds" in inspect.signature(create_causal_mask).parameters else "input_embeds"
16
+
17
+
18
+ class GPT2AttentionModified(GPT2Attention):
19
+ def forward(
20
+ self,
21
+ hidden_states: Optional[tuple[torch.FloatTensor]],
22
+ past_key_values: Optional[Cache] = None,
23
+ cache_position: Optional[torch.LongTensor] = None,
24
+ attention_mask: Optional[torch.FloatTensor] = None,
25
+ head_mask: Optional[torch.FloatTensor] = None,
26
+ encoder_hidden_states: Optional[torch.Tensor] = None,
27
+ encoder_attention_mask: Optional[torch.FloatTensor] = None,
28
+ output_attentions: Optional[bool] = False,
29
+ **kwargs,
30
+ ):
31
+ is_cross_attention = encoder_hidden_states is not None
32
+ if past_key_values is not None:
33
+ if isinstance(past_key_values, EncoderDecoderCache):
34
+ is_updated = past_key_values.is_updated.get(self.layer_idx)
35
+ curr_past_key_value = past_key_values.cross_attention_cache if is_cross_attention else past_key_values.self_attention_cache
36
+ else:
37
+ curr_past_key_value = past_key_values
38
+
39
+ if is_cross_attention:
40
+ if not hasattr(self, "q_attn"):
41
+ raise ValueError("Cross-attention requires q_attn to be defined.")
42
+ query_states = self.q_attn(hidden_states)
43
+ attention_mask = encoder_attention_mask
44
+ if past_key_values is not None and is_updated:
45
+ key_states = curr_past_key_value.layers[self.layer_idx].keys
46
+ value_states = curr_past_key_value.layers[self.layer_idx].values
47
+ else:
48
+ key_states, value_states = self.c_attn(encoder_hidden_states).split(self.split_size, dim=2)
49
+ shape_kv = (*key_states.shape[:-1], -1, self.head_dim)
50
+ key_states = key_states.view(shape_kv).transpose(1, 2)
51
+ value_states = value_states.view(shape_kv).transpose(1, 2)
52
+ else:
53
+ query_states, key_states, value_states = self.c_attn(hidden_states).split(self.split_size, dim=2)
54
+ shape_kv = (*key_states.shape[:-1], -1, self.head_dim)
55
+ key_states = key_states.view(shape_kv).transpose(1, 2)
56
+ value_states = value_states.view(shape_kv).transpose(1, 2)
57
+
58
+ shape_q = (*query_states.shape[:-1], -1, self.head_dim)
59
+ query_states = query_states.view(shape_q).transpose(1, 2)
60
+
61
+ if (past_key_values is not None and not is_cross_attention) or (
62
+ past_key_values is not None and is_cross_attention and not is_updated
63
+ ):
64
+ cache_position = cache_position if not is_cross_attention else None
65
+ key_states, value_states = curr_past_key_value.update(
66
+ key_states, value_states, self.layer_idx, {"cache_position": cache_position}
67
+ )
68
+ if is_cross_attention:
69
+ past_key_values.is_updated[self.layer_idx] = True
70
+
71
+ is_causal = attention_mask is None and query_states.shape[-2] > 1 and not is_cross_attention
72
+ attention_interface = eager_attention_forward
73
+ if self.config._attn_implementation != "eager":
74
+ attention_interface = ALL_ATTENTION_FUNCTIONS[self.config._attn_implementation]
75
+
76
+ attn_output, attn_weights = attention_interface(
77
+ self,
78
+ query_states,
79
+ key_states,
80
+ value_states,
81
+ attention_mask,
82
+ head_mask=head_mask,
83
+ dropout=self.attn_dropout.p if self.training else 0.0,
84
+ is_causal=is_causal,
85
+ **kwargs,
86
+ )
87
+
88
+ attn_output = attn_output.reshape(*attn_output.shape[:-2], -1).contiguous()
89
+ attn_output = self.c_proj(attn_output)
90
+ attn_output = self.resid_dropout(attn_output)
91
+ return attn_output, attn_weights
92
+
93
+
94
+ class GPT2BlockModified(GPT2Block):
95
+ def __init__(self, config, layer_idx=None):
96
+ super().__init__(config=config, layer_idx=layer_idx)
97
+ self.attn = GPT2AttentionModified(config=config, layer_idx=layer_idx)
98
+
99
+
100
+ class GPT2ModelModified(GPT2Model):
101
+ def __init__(self, config):
102
+ super().__init__(config)
103
+ self.config_causal = config
104
+ self.config_causal._attn_implementation = "eager"
105
+ self.h = nn.ModuleList([GPT2BlockModified(config, layer_idx=i) for i in range(config.num_hidden_layers)])
106
+
107
+ def forward(
108
+ self,
109
+ input_ids: Optional[torch.LongTensor] = None,
110
+ past_key_values: Optional[Union[tuple[tuple[torch.Tensor]], Cache]] = None,
111
+ cache_position: Optional[torch.LongTensor] = None,
112
+ attention_mask: Optional[torch.FloatTensor] = None,
113
+ token_type_ids: Optional[torch.LongTensor] = None,
114
+ position_ids: Optional[torch.LongTensor] = None,
115
+ head_mask: Optional[torch.FloatTensor] = None,
116
+ inputs_embeds: Optional[torch.FloatTensor] = None,
117
+ encoder_hidden_states: Optional[torch.Tensor] = None,
118
+ encoder_attention_mask: Optional[torch.FloatTensor] = None,
119
+ use_cache: Optional[bool] = None,
120
+ output_attentions: Optional[bool] = None,
121
+ output_hidden_states: Optional[bool] = None,
122
+ return_dict: Optional[bool] = None,
123
+ segmentation_mask: Optional[torch.FloatTensor] = None,
124
+ **kwargs,
125
+ ) -> Union[tuple, BaseModelOutputWithPastAndCrossAttentions]:
126
+ output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions
127
+ output_hidden_states = output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
128
+ use_cache = use_cache if use_cache is not None else self.config.use_cache
129
+ return_dict = return_dict if return_dict is not None else self.config.use_return_dict
130
+
131
+ if input_ids is not None and inputs_embeds is not None:
132
+ raise ValueError("You cannot specify both input_ids and inputs_embeds at the same time")
133
+ if input_ids is not None:
134
+ self.warn_if_padding_and_no_attention_mask(input_ids, attention_mask)
135
+ input_shape = input_ids.size()
136
+ input_ids = input_ids.view(-1, input_shape[-1])
137
+ batch_size = input_ids.shape[0]
138
+ elif inputs_embeds is not None:
139
+ input_shape = inputs_embeds.size()[:-1]
140
+ batch_size = inputs_embeds.shape[0]
141
+ else:
142
+ raise ValueError("You have to specify either input_ids or inputs_embeds")
143
+
144
+ device = input_ids.device if input_ids is not None else inputs_embeds.device
145
+
146
+ if token_type_ids is not None:
147
+ token_type_ids = token_type_ids.view(-1, input_shape[-1])
148
+
149
+ if self.gradient_checkpointing and self.training and use_cache:
150
+ use_cache = False
151
+
152
+ if use_cache:
153
+ if past_key_values is None:
154
+ past_key_values = DynamicCache()
155
+ elif isinstance(past_key_values, tuple):
156
+ past_key_values = DynamicCache.from_legacy_cache(past_key_values)
157
+ if self.config.add_cross_attention and not isinstance(past_key_values, EncoderDecoderCache):
158
+ past_key_values = EncoderDecoderCache(past_key_values, DynamicCache())
159
+
160
+ if inputs_embeds is None:
161
+ inputs_embeds = self.wte(input_ids)
162
+
163
+ if cache_position is None:
164
+ past_seen_tokens = past_key_values.get_seq_length() if past_key_values is not None else 0
165
+ cache_position = torch.arange(past_seen_tokens, past_seen_tokens + inputs_embeds.shape[1], device=inputs_embeds.device)
166
+ if position_ids is None:
167
+ position_ids = cache_position.unsqueeze(0)
168
+
169
+ position_embeds = self.wpe(position_ids)
170
+ hidden_states = inputs_embeds + position_embeds.to(inputs_embeds.device)
171
+
172
+ if attention_mask is not None and attention_mask.ndim < 4:
173
+ attention_mask = attention_mask.view(batch_size, -1)
174
+
175
+ causal_mask_kwargs = {
176
+ "config": self.config_causal,
177
+ _CREATE_CAUSAL_MASK_EMBEDS_ARG: inputs_embeds,
178
+ "attention_mask": attention_mask,
179
+ "cache_position": cache_position,
180
+ "past_key_values": past_key_values,
181
+ "position_ids": position_ids,
182
+ }
183
+ causal_mask = create_causal_mask(**causal_mask_kwargs)
184
+
185
+ _use_sdpa = self._attn_implementation == "sdpa" and output_attentions is False and head_mask is None
186
+ if self.config.add_cross_attention and encoder_hidden_states is not None:
187
+ encoder_batch_size, encoder_sequence_length, _ = encoder_hidden_states.size()
188
+ encoder_hidden_shape = (encoder_batch_size, encoder_sequence_length)
189
+ if encoder_attention_mask is None:
190
+ encoder_attention_mask = torch.ones(encoder_hidden_shape, device=device)
191
+ if _use_sdpa:
192
+ encoder_attention_mask = _prepare_4d_attention_mask_for_sdpa(
193
+ mask=encoder_attention_mask, dtype=inputs_embeds.dtype, tgt_len=input_shape[-1]
194
+ )
195
+ elif self._attn_implementation != "flash_attention_2":
196
+ encoder_attention_mask = self.invert_attention_mask(encoder_attention_mask)
197
+ else:
198
+ encoder_attention_mask = None
199
+
200
+ if head_mask is None:
201
+ head_mask = [None] * self.config.n_layer
202
+
203
+ if token_type_ids is not None:
204
+ hidden_states = hidden_states + self.wte(token_type_ids)
205
+
206
+ hidden_states = self.drop(hidden_states)
207
+ output_shape = (-1,) + input_shape[1:] + (hidden_states.size(-1),)
208
+ all_self_attentions = () if output_attentions else None
209
+ all_cross_attentions = () if output_attentions and self.config.add_cross_attention else None
210
+ all_hidden_states = () if output_hidden_states else None
211
+
212
+ for i, block in enumerate(self.h):
213
+ if output_hidden_states:
214
+ all_hidden_states = all_hidden_states + (hidden_states,)
215
+
216
+ block_mask = causal_mask
217
+ if segmentation_mask is not None and causal_mask is not None:
218
+ block_mask = causal_mask.clone()
219
+ seq_len = input_shape[-1]
220
+ if block_mask.shape[2] != seq_len or block_mask.shape[3] != seq_len:
221
+ block_mask = block_mask[:, :, :seq_len, :seq_len]
222
+ layer_bias = segmentation_mask[:, i, : block_mask.shape[2], : block_mask.shape[3]].unsqueeze(1)
223
+ block_mask = block_mask + layer_bias.to(dtype=block_mask.dtype, device=block_mask.device)
224
+
225
+ outputs = block(
226
+ hidden_states=hidden_states,
227
+ past_key_values=past_key_values if not (self.gradient_checkpointing and self.training) else None,
228
+ cache_position=cache_position,
229
+ attention_mask=block_mask,
230
+ encoder_hidden_states=encoder_hidden_states,
231
+ encoder_attention_mask=encoder_attention_mask,
232
+ use_cache=use_cache,
233
+ output_attentions=output_attentions,
234
+ head_mask=head_mask[i],
235
+ **kwargs,
236
+ )
237
+ if isinstance(outputs, tuple):
238
+ hidden_states = outputs[0]
239
+ if output_attentions and len(outputs) > 1:
240
+ all_self_attentions = all_self_attentions + (outputs[1],)
241
+ if self.config.add_cross_attention and len(outputs) > 2:
242
+ all_cross_attentions = all_cross_attentions + (outputs[2],)
243
+ else:
244
+ hidden_states = outputs
245
+
246
+ hidden_states = self.ln_f(hidden_states)
247
+ hidden_states = hidden_states.view(output_shape)
248
+ if output_hidden_states:
249
+ all_hidden_states = all_hidden_states + (hidden_states,)
250
+
251
+ past_key_values = past_key_values if use_cache else None
252
+ if not return_dict:
253
+ return tuple(v for v in [hidden_states, past_key_values, all_hidden_states, all_self_attentions, all_cross_attentions] if v is not None)
254
+
255
+ return BaseModelOutputWithPastAndCrossAttentions(
256
+ last_hidden_state=hidden_states,
257
+ past_key_values=past_key_values,
258
+ hidden_states=all_hidden_states,
259
+ attentions=all_self_attentions,
260
+ cross_attentions=all_cross_attentions,
261
+ )
262
+
263
+
264
+ class GPT2LMHeadModelModified(GPT2LMHeadModel):
265
+ def __init__(self, config):
266
+ super().__init__(config)
267
+ self.transformer = GPT2ModelModified(config)
268
+ self.post_init()
269
+
270
+ def forward(
271
+ self,
272
+ input_ids: Optional[torch.LongTensor] = None,
273
+ past_key_values: Optional[tuple[tuple[torch.Tensor]]] = None,
274
+ cache_position: Optional[torch.LongTensor] = None,
275
+ attention_mask: Optional[torch.FloatTensor] = None,
276
+ token_type_ids: Optional[torch.LongTensor] = None,
277
+ position_ids: Optional[torch.LongTensor] = None,
278
+ head_mask: Optional[torch.FloatTensor] = None,
279
+ inputs_embeds: Optional[torch.FloatTensor] = None,
280
+ encoder_hidden_states: Optional[torch.Tensor] = None,
281
+ encoder_attention_mask: Optional[torch.FloatTensor] = None,
282
+ labels: Optional[torch.LongTensor] = None,
283
+ use_cache: Optional[bool] = None,
284
+ output_attentions: Optional[bool] = None,
285
+ output_hidden_states: Optional[bool] = None,
286
+ return_dict: Optional[bool] = None,
287
+ logits_to_keep: Union[int, torch.Tensor] = 0,
288
+ segmentation_mask: Optional[torch.FloatTensor] = None,
289
+ **kwargs,
290
+ ) -> Union[tuple, CausalLMOutputWithCrossAttentions]:
291
+ return_dict = return_dict if return_dict is not None else self.config.use_return_dict
292
+ transformer_outputs = self.transformer(
293
+ input_ids,
294
+ past_key_values=past_key_values,
295
+ attention_mask=attention_mask,
296
+ cache_position=cache_position,
297
+ token_type_ids=token_type_ids,
298
+ position_ids=position_ids,
299
+ head_mask=head_mask,
300
+ inputs_embeds=inputs_embeds,
301
+ encoder_hidden_states=encoder_hidden_states,
302
+ encoder_attention_mask=encoder_attention_mask,
303
+ use_cache=use_cache,
304
+ output_attentions=output_attentions,
305
+ output_hidden_states=output_hidden_states,
306
+ return_dict=return_dict,
307
+ segmentation_mask=segmentation_mask,
308
+ **kwargs,
309
+ )
310
+ hidden_states = transformer_outputs[0]
311
+ slice_indices = slice(-logits_to_keep, None) if isinstance(logits_to_keep, int) and logits_to_keep > 0 else slice(None)
312
+ logits = self.lm_head(hidden_states[:, slice_indices, :])
313
+
314
+ loss = None
315
+ if labels is not None:
316
+ loss = self.loss_function(logits, labels, vocab_size=self.config.vocab_size, **kwargs)
317
+
318
+ if not return_dict:
319
+ output = (logits,) + transformer_outputs[1:]
320
+ return ((loss,) + output) if loss is not None else output
321
+
322
+ return CausalLMOutputWithCrossAttentions(
323
+ loss=loss,
324
+ logits=logits,
325
+ past_key_values=transformer_outputs.past_key_values,
326
+ hidden_states=transformer_outputs.hidden_states,
327
+ attentions=transformer_outputs.attentions,
328
+ cross_attentions=transformer_outputs.cross_attentions,
329
+ )
330
+
331
+
332
+ @torch.no_grad()
333
+ def expand_gpt2_positional_embeddings(
334
+ model: torch.nn.Module,
335
+ new_max_positions: int,
336
+ mode: str = "linear",
337
+ align_corners: bool = True,
338
+ ):
339
+ if hasattr(model, "transformer") and hasattr(model.transformer, "wpe"):
340
+ model_for_wpe = model.transformer
341
+ elif hasattr(model, "wpe"):
342
+ model_for_wpe = model
343
+ else:
344
+ raise ValueError("Model does not expose GPT-2 positional embeddings.")
345
+
346
+ wpe = model_for_wpe.wpe
347
+ old_n, d = wpe.weight.shape
348
+ if new_max_positions == old_n:
349
+ return model
350
+
351
+ device = wpe.weight.device
352
+ dtype = wpe.weight.dtype
353
+ if new_max_positions < old_n:
354
+ new_weight = wpe.weight[:new_max_positions].clone()
355
+ else:
356
+ if mode != "linear":
357
+ raise ValueError(f"Unsupported positional expansion mode: {mode}")
358
+ w = wpe.weight.transpose(0, 1).unsqueeze(0)
359
+ w_new = F.interpolate(w, size=new_max_positions, mode="linear", align_corners=align_corners)
360
+ new_weight = w_new.squeeze(0).transpose(0, 1).contiguous()
361
+
362
+ new_wpe = torch.nn.Embedding(new_max_positions, d, device=device, dtype=dtype)
363
+ new_wpe.weight.copy_(new_weight)
364
+ if hasattr(model, "transformer") and hasattr(model.transformer, "wpe"):
365
+ model.transformer.wpe = new_wpe
366
+ else:
367
+ model.wpe = new_wpe
368
+ if hasattr(model.config, "n_positions"):
369
+ model.config.n_positions = new_max_positions
370
+ if hasattr(model.config, "n_ctx"):
371
+ model.config.n_ctx = new_max_positions
372
+ return model
373
+
374
+
375
+ def create_decoder(
376
+ text_model_name: str,
377
+ attention_implementation: str,
378
+ max_position_embeddings: int,
379
+ load_pretrained: bool = True,
380
+ vocab_size: Optional[int] = None,
381
+ pad_token_id: Optional[int] = None,
382
+ **decoder_kwargs,
383
+ ):
384
+ config = GPT2Config.from_pretrained(text_model_name)
385
+ config._attn_implementation = attention_implementation
386
+ config.n_positions = max_position_embeddings
387
+ config.n_ctx = max_position_embeddings
388
+ config.tie_word_embeddings = False
389
+ if vocab_size is not None:
390
+ config.vocab_size = vocab_size
391
+ if pad_token_id is not None:
392
+ config.pad_token_id = pad_token_id
393
+ config.use_cache = decoder_kwargs.pop("use_cache", True)
394
+ if load_pretrained:
395
+ decoder = GPT2LMHeadModelModified.from_pretrained(text_model_name, config=config, **decoder_kwargs)
396
+ else:
397
+ decoder = GPT2LMHeadModelModified(config)
398
+ decoder.config._attn_implementation = attention_implementation
399
+ return expand_gpt2_positional_embeddings(decoder, new_max_positions=max_position_embeddings, mode="linear")
image_processing_lana.py ADDED
@@ -0,0 +1,85 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ from typing import Any
4
+
5
+ import numpy as np
6
+ from transformers.image_processing_utils import BaseImageProcessor, BatchFeature, get_size_dict
7
+ from transformers.image_transforms import convert_to_rgb, normalize, resize, to_channel_dimension_format
8
+ from transformers.image_utils import (
9
+ ChannelDimension,
10
+ ImageInput,
11
+ PILImageResampling,
12
+ infer_channel_dimension_format,
13
+ make_flat_list_of_images,
14
+ to_numpy_array,
15
+ valid_images,
16
+ )
17
+ from transformers.utils import TensorType
18
+
19
+
20
+ class LanaImageProcessor(BaseImageProcessor):
21
+ model_input_names = ["pixel_values"]
22
+
23
+ def __init__(
24
+ self,
25
+ do_resize: bool = True,
26
+ size: dict[str, int] | None = None,
27
+ resample: PILImageResampling = PILImageResampling.BICUBIC,
28
+ do_rescale: bool = True,
29
+ rescale_factor: float = 1 / 255.0,
30
+ do_normalize: bool = True,
31
+ image_mean: list[float] | None = None,
32
+ image_std: list[float] | None = None,
33
+ do_convert_rgb: bool = True,
34
+ **kwargs,
35
+ ) -> None:
36
+ super().__init__(**kwargs)
37
+ self.do_resize = do_resize
38
+ self.size = get_size_dict(size or {"height": 512, "width": 512})
39
+ self.resample = resample
40
+ self.do_rescale = do_rescale
41
+ self.rescale_factor = rescale_factor
42
+ self.do_normalize = do_normalize
43
+ self.image_mean = image_mean or [0.485, 0.456, 0.406]
44
+ self.image_std = image_std or [0.229, 0.224, 0.225]
45
+ self.do_convert_rgb = do_convert_rgb
46
+
47
+ def preprocess(
48
+ self,
49
+ images: ImageInput,
50
+ return_tensors: str | TensorType | None = None,
51
+ data_format: ChannelDimension = ChannelDimension.FIRST,
52
+ **kwargs: Any,
53
+ ) -> BatchFeature:
54
+ images = make_flat_list_of_images(images)
55
+ if not valid_images(images):
56
+ raise ValueError("LanaImageProcessor expected a PIL image, numpy array, torch tensor, or a list of images.")
57
+
58
+ pixel_values = []
59
+ for image in images:
60
+ if self.do_convert_rgb:
61
+ image = convert_to_rgb(image)
62
+ array = to_numpy_array(image).astype(np.float32)
63
+ input_data_format = infer_channel_dimension_format(array)
64
+ if self.do_resize:
65
+ array = resize(
66
+ image=array,
67
+ size=(self.size["height"], self.size["width"]),
68
+ resample=self.resample,
69
+ input_data_format=input_data_format,
70
+ )
71
+ input_data_format = infer_channel_dimension_format(array)
72
+ if self.do_rescale:
73
+ array = array * self.rescale_factor
74
+ if self.do_normalize:
75
+ array = normalize(
76
+ array,
77
+ mean=self.image_mean,
78
+ std=self.image_std,
79
+ input_data_format=input_data_format,
80
+ )
81
+ array = to_channel_dimension_format(array, data_format, input_channel_dim=input_data_format)
82
+ array = np.asarray(array, dtype=np.float32)
83
+ pixel_values.append(array)
84
+
85
+ return BatchFeature(data={"pixel_values": pixel_values}, tensor_type=return_tensors)
layerwise_anatomical_attention.py ADDED
@@ -0,0 +1,122 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+ import torch.nn.functional as F
3
+
4
+
5
+ def _gaussian_kernel_1d(kernel_size: int, sigma: float, device: torch.device, dtype: torch.dtype) -> torch.Tensor:
6
+ radius = kernel_size // 2
7
+ x = torch.arange(-radius, radius + 1, device=device, dtype=dtype)
8
+ kernel = torch.exp(-(x * x) / (2.0 * sigma * sigma))
9
+ return kernel / kernel.sum()
10
+
11
+
12
+ @torch.no_grad()
13
+ def build_layerwise_attention_bias(
14
+ masks: torch.Tensor,
15
+ num_layers: int,
16
+ target_tokens: int,
17
+ base_kernel_size: int = 3,
18
+ kernel_growth: int = 2,
19
+ strength: float = 2.0,
20
+ eps: float = 1e-8,
21
+ ) -> torch.Tensor:
22
+ if masks.ndim == 3:
23
+ masks = masks.unsqueeze(1)
24
+ if masks.ndim != 4 or masks.shape[1] != 1:
25
+ raise ValueError(f"Expected masks shaped (B,1,H,W) or (B,H,W), got {tuple(masks.shape)}")
26
+
27
+ masks = masks.float()
28
+ batch_size = masks.shape[0]
29
+ resized = F.interpolate(masks, size=(32, 32), mode="bilinear", align_corners=False).clamp(0.0, 1.0)
30
+
31
+ max_kernel = base_kernel_size + max(num_layers, 0) * kernel_growth
32
+ if max_kernel % 2 == 0:
33
+ max_kernel += 1
34
+ pad = max_kernel // 2
35
+
36
+ weight_h = torch.zeros((num_layers, 1, 1, max_kernel), device=resized.device, dtype=resized.dtype)
37
+ weight_v = torch.zeros((num_layers, 1, max_kernel, 1), device=resized.device, dtype=resized.dtype)
38
+
39
+ for layer_idx in range(num_layers):
40
+ kernel_size = base_kernel_size + (num_layers - layer_idx) * kernel_growth
41
+ if kernel_size % 2 == 0:
42
+ kernel_size += 1
43
+ sigma = max((kernel_size - 1) / 6.0, 1e-3)
44
+ kernel = _gaussian_kernel_1d(kernel_size, sigma, resized.device, resized.dtype)
45
+ start = (max_kernel - kernel_size) // 2
46
+ end = start + kernel_size
47
+ weight_h[layer_idx, 0, 0, start:end] = kernel
48
+ weight_v[layer_idx, 0, start:end, 0] = kernel
49
+
50
+ repeated = resized.expand(batch_size, num_layers, 32, 32).contiguous()
51
+ horizontal = F.conv2d(F.pad(repeated, (pad, pad, 0, 0), mode="reflect"), weight_h, groups=num_layers)
52
+ vertical = F.conv2d(F.pad(horizontal, (0, 0, pad, pad), mode="reflect"), weight_v, groups=num_layers)
53
+
54
+ min_vals = vertical.amin(dim=(2, 3), keepdim=True)
55
+ max_vals = vertical.amax(dim=(2, 3), keepdim=True)
56
+ normalized = (vertical - min_vals) / (max_vals - min_vals).clamp_min(eps)
57
+
58
+ flat = normalized.view(batch_size, num_layers, -1)
59
+ if flat.shape[-1] != target_tokens:
60
+ flat = F.interpolate(flat, size=target_tokens, mode="linear", align_corners=False)
61
+ layerwise_bias = flat.unsqueeze(-2).expand(-1, -1, target_tokens, -1)
62
+ return torch.tril(layerwise_bias) * strength
63
+
64
+
65
+ @torch.no_grad()
66
+ def build_legacy_gaussian_attention_bias(
67
+ masks: torch.Tensor,
68
+ num_layers: int,
69
+ target_query_tokens: int,
70
+ target_key_tokens: int,
71
+ base_kernel_size: int = 3,
72
+ kernel_growth: int = 2,
73
+ strength: float = 1.0,
74
+ eps: float = 1e-8,
75
+ ) -> torch.Tensor:
76
+ if masks.ndim == 3:
77
+ masks = masks.unsqueeze(1)
78
+ if masks.ndim != 4 or masks.shape[1] != 1:
79
+ raise ValueError(f"Expected masks shaped (B,1,H,W) or (B,H,W), got {tuple(masks.shape)}")
80
+
81
+ masks = masks.float()
82
+ batch_size = masks.shape[0]
83
+ xmin = masks.amin(dim=(2, 3), keepdim=True)
84
+ xmax = masks.amax(dim=(2, 3), keepdim=True)
85
+ normalized_masks = (masks - xmin) / (xmax - xmin).clamp_min(eps)
86
+ resized = F.interpolate(normalized_masks, size=(32, 32), mode="bilinear", align_corners=False)
87
+
88
+ kernel_sizes = []
89
+ for layer_idx in range(num_layers, 0, -1):
90
+ kernel_size = base_kernel_size + layer_idx * kernel_growth
91
+ if kernel_size % 2 == 0:
92
+ kernel_size += 1
93
+ kernel_sizes.append(max(kernel_size, 1))
94
+
95
+ max_kernel = max(kernel_sizes)
96
+ pad = max_kernel // 2
97
+ weight_h = torch.zeros((num_layers, 1, 1, max_kernel), device=resized.device, dtype=resized.dtype)
98
+ weight_v = torch.zeros((num_layers, 1, max_kernel, 1), device=resized.device, dtype=resized.dtype)
99
+
100
+ for layer_idx, kernel_size in enumerate(kernel_sizes):
101
+ sigma = max((kernel_size - 1) / 6.0, 1e-3)
102
+ kernel = _gaussian_kernel_1d(kernel_size, sigma, resized.device, resized.dtype)
103
+ start = (max_kernel - kernel_size) // 2
104
+ end = start + kernel_size
105
+ weight_h[layer_idx, 0, 0, start:end] = kernel
106
+ weight_v[layer_idx, 0, start:end, 0] = kernel
107
+
108
+ repeated = resized.expand(batch_size, num_layers, 32, 32).contiguous()
109
+ horizontal = F.conv2d(F.pad(repeated, (pad, pad, 0, 0), mode="reflect"), weight_h, groups=num_layers)
110
+ vertical = F.conv2d(F.pad(horizontal, (0, 0, pad, pad), mode="reflect"), weight_v, groups=num_layers)
111
+
112
+ min_vals = vertical.amin(dim=(2, 3), keepdim=True)
113
+ max_vals = vertical.amax(dim=(2, 3), keepdim=True)
114
+ normalized = (vertical - min_vals) / (max_vals - min_vals).clamp_min(eps)
115
+
116
+ flat = normalized.view(batch_size, num_layers, -1)
117
+ if flat.shape[-1] != target_key_tokens:
118
+ flat = F.interpolate(flat, size=target_key_tokens, mode="linear", align_corners=False)
119
+ return flat.unsqueeze(-2).expand(-1, -1, target_query_tokens, -1) * strength
120
+
121
+
122
+ __all__ = ["build_layerwise_attention_bias", "build_legacy_gaussian_attention_bias"]
merges.txt ADDED
The diff for this file is too large to render. See raw diff
 
model.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:008ac4050299165c2523c84179c58c9bdbec82096208a64d10be01b955115109
3
+ size 1152540320
modeling_lana.py ADDED
@@ -0,0 +1,393 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import logging
2
+ from pathlib import Path
3
+ from typing import Optional
4
+
5
+ import torch
6
+ import torch.nn as nn
7
+ from huggingface_hub import snapshot_download
8
+ from transformers import AutoConfig, AutoModel, AutoTokenizer, BitsAndBytesConfig, GPT2Tokenizer, PreTrainedModel
9
+
10
+ from .configuration_lana import LanaConfig
11
+ from .gpt2_modified import create_decoder
12
+ from .layerwise_anatomical_attention import build_layerwise_attention_bias, build_legacy_gaussian_attention_bias
13
+ from .modeling_outputs import LanaModelOutput
14
+ from .segmenters import AnatomicalSegmenter
15
+
16
+ logger = logging.getLogger(__name__)
17
+ PAD_TOKEN = "<|pad|>"
18
+
19
+
20
+ def _resolve_repo_root(config: LanaConfig) -> Path | None:
21
+ for candidate in [getattr(config, "local_repo_path", ""), getattr(config, "_name_or_path", "")]:
22
+ if not candidate:
23
+ continue
24
+ path = Path(str(candidate))
25
+ if path.exists():
26
+ return path
27
+ return None
28
+
29
+
30
+ def _resolve_source(reference: str, repo_root: Path | None) -> str:
31
+ if not reference:
32
+ return reference
33
+ path = Path(reference)
34
+ if path.is_absolute() and path.exists():
35
+ return str(path)
36
+ if repo_root is not None:
37
+ repo_path = repo_root / reference
38
+ if repo_path.exists():
39
+ return str(repo_path)
40
+ if path.exists():
41
+ return str(path)
42
+ return reference
43
+
44
+
45
+ def _resolve_tokenizer_source(config: LanaConfig, repo_root: Path | None) -> str:
46
+ for reference in [
47
+ getattr(config, "bundled_tokenizer_name", ""),
48
+ "",
49
+ ]:
50
+ if reference:
51
+ resolved = _resolve_source(reference, repo_root)
52
+ if resolved and Path(resolved).exists():
53
+ return resolved
54
+ if repo_root is not None and (repo_root / "tokenizer_config.json").exists():
55
+ return str(repo_root)
56
+ return _resolve_source(config.text_model_name, repo_root)
57
+
58
+
59
+ def _is_local_source(reference: str, repo_root: Path | None) -> bool:
60
+ resolved = _resolve_source(reference, repo_root)
61
+ return bool(resolved) and Path(resolved).exists()
62
+
63
+
64
+ def build_visual_projection(config: LanaConfig) -> nn.Module:
65
+ if config.visual_projection_type == "linear":
66
+ return nn.Linear(config.visual_feature_dim, config.text_hidden_size)
67
+ if config.visual_projection_type == "mlp4":
68
+ return nn.Sequential(
69
+ nn.Linear(config.visual_feature_dim, config.text_hidden_size),
70
+ nn.GELU(),
71
+ nn.Linear(config.text_hidden_size, config.text_hidden_size),
72
+ nn.GELU(),
73
+ nn.Linear(config.text_hidden_size, config.text_hidden_size),
74
+ nn.GELU(),
75
+ nn.Linear(config.text_hidden_size, config.text_hidden_size),
76
+ )
77
+ raise ValueError(f"Unsupported visual projection type: {config.visual_projection_type}")
78
+
79
+
80
+ class LanaForConditionalGeneration(PreTrainedModel):
81
+ config_class = LanaConfig
82
+ base_model_prefix = "lana"
83
+ supports_gradient_checkpointing = True
84
+
85
+ def __init__(self, config: LanaConfig):
86
+ super().__init__(config)
87
+ repo_root = _resolve_repo_root(config)
88
+ vision_model_name = _resolve_source(getattr(config, "bundled_vision_model_name", "") or config.vision_model_name, repo_root)
89
+ text_model_name = _resolve_source(getattr(config, "bundled_text_model_name", "") or config.text_model_name, repo_root)
90
+ segmentation_model_name = _resolve_source(
91
+ getattr(config, "bundled_segmentation_model_name", "") or config.segmentation_model_name,
92
+ repo_root,
93
+ )
94
+ tokenizer_source = _resolve_tokenizer_source(config, repo_root)
95
+ lung_checkpoint = _resolve_source(config.lung_segmenter_checkpoint, repo_root)
96
+ heart_checkpoint = _resolve_source(config.heart_segmenter_checkpoint, repo_root)
97
+ segmenter_weights_in_model_state = bool(getattr(config, "segmenter_weights_in_model_state", False))
98
+
99
+ vision_config = AutoConfig.from_pretrained(vision_model_name, trust_remote_code=True)
100
+ if getattr(vision_config, "hidden_size", None) is not None:
101
+ config.visual_feature_dim = vision_config.hidden_size
102
+
103
+ vision_load_pretrained = not _is_local_source(vision_model_name, repo_root)
104
+ if vision_load_pretrained:
105
+ self.vision_encoder = AutoModel.from_pretrained(vision_model_name, trust_remote_code=True)
106
+ else:
107
+ self.vision_encoder = AutoModel.from_config(vision_config, trust_remote_code=True)
108
+ decoder_kwargs = {
109
+ "ignore_mismatched_sizes": True,
110
+ "use_cache": config.use_cache,
111
+ }
112
+ if config.decoder_load_in_4bit:
113
+ compute_dtype = getattr(torch, config.decoder_compute_dtype, torch.float16)
114
+ decoder_kwargs["quantization_config"] = BitsAndBytesConfig(
115
+ load_in_4bit=True,
116
+ bnb_4bit_quant_type="nf4",
117
+ bnb_4bit_use_double_quant=True,
118
+ bnb_4bit_compute_dtype=compute_dtype,
119
+ )
120
+ decoder_kwargs["device_map"] = {"": 0}
121
+ self.text_decoder = create_decoder(
122
+ text_model_name=text_model_name,
123
+ attention_implementation=config.segmentation_attention_implementation,
124
+ max_position_embeddings=config.max_position_embeddings,
125
+ load_pretrained=not _is_local_source(text_model_name, repo_root),
126
+ vocab_size=getattr(config, "vocab_size", None),
127
+ **decoder_kwargs,
128
+ )
129
+ if _is_local_source(tokenizer_source, repo_root):
130
+ self.tokenizer = GPT2Tokenizer.from_pretrained(tokenizer_source)
131
+ else:
132
+ self.tokenizer = AutoTokenizer.from_pretrained(tokenizer_source, trust_remote_code=True, use_fast=False)
133
+ if self.tokenizer.pad_token_id is None:
134
+ target_vocab_size = getattr(config, "vocab_size", None)
135
+ if target_vocab_size and target_vocab_size > len(self.tokenizer):
136
+ self.tokenizer.add_special_tokens({"pad_token": PAD_TOKEN})
137
+ else:
138
+ fallback_pad = self.tokenizer.eos_token or self.tokenizer.bos_token or PAD_TOKEN
139
+ self.tokenizer.pad_token = fallback_pad
140
+ if self.text_decoder.get_input_embeddings().weight.shape[0] != len(self.tokenizer):
141
+ self.text_decoder.resize_token_embeddings(len(self.tokenizer))
142
+ self.text_decoder.config.pad_token_id = self.tokenizer.pad_token_id
143
+ if hasattr(self.text_decoder, "generation_config") and self.text_decoder.generation_config is not None:
144
+ self.text_decoder.generation_config.pad_token_id = self.tokenizer.pad_token_id
145
+ self.text_decoder.generation_config.eos_token_id = None
146
+
147
+ config.vocab_size = self.text_decoder.config.vocab_size
148
+ config.text_hidden_size = self.text_decoder.config.hidden_size
149
+ config.num_attention_layers = self.text_decoder.config.n_layer
150
+
151
+ self.visual_projection = build_visual_projection(config)
152
+ self.segmenter = None
153
+ if config.use_segmentation_mask:
154
+ assume_segmenter_weights_from_model_state = segmenter_weights_in_model_state and not (
155
+ Path(lung_checkpoint).exists() or Path(heart_checkpoint).exists()
156
+ )
157
+ self.segmenter = AnatomicalSegmenter(
158
+ model_name=segmentation_model_name,
159
+ freeze=config.freeze_segmenter,
160
+ lung_checkpoint=lung_checkpoint,
161
+ heart_checkpoint=heart_checkpoint,
162
+ load_pretrained=not _is_local_source(segmentation_model_name, repo_root),
163
+ assume_weights_from_model_state=assume_segmenter_weights_from_model_state,
164
+ )
165
+ self.post_init()
166
+
167
+ @classmethod
168
+ def from_pretrained(cls, pretrained_model_name_or_path, *model_args, **kwargs):
169
+ kwargs.setdefault("low_cpu_mem_usage", False)
170
+ config = kwargs.get("config")
171
+ if config is not None and getattr(config, "local_repo_path", ""):
172
+ return super().from_pretrained(pretrained_model_name_or_path, *model_args, **kwargs)
173
+
174
+ repo_path = str(pretrained_model_name_or_path)
175
+ if not Path(repo_path).exists():
176
+ repo_path = snapshot_download(repo_path)
177
+
178
+ if config is None:
179
+ config = LanaConfig.from_pretrained(repo_path, trust_remote_code=True)
180
+ config.local_repo_path = repo_path
181
+ kwargs["config"] = config
182
+ return super().from_pretrained(repo_path, *model_args, **kwargs)
183
+
184
+ def move_non_quantized_modules(self, device: torch.device) -> None:
185
+ self.vision_encoder.to(device)
186
+ self.visual_projection.to(device)
187
+ if self.segmenter is not None:
188
+ self.segmenter.to(device)
189
+ if not getattr(self.config, "decoder_load_in_4bit", False):
190
+ self.text_decoder.to(device)
191
+
192
+ def _encode_images(self, pixel_values: torch.Tensor) -> torch.Tensor:
193
+ if any(param.requires_grad for param in self.vision_encoder.parameters()):
194
+ outputs = self.vision_encoder(pixel_values=pixel_values)
195
+ else:
196
+ with torch.no_grad():
197
+ outputs = self.vision_encoder(pixel_values=pixel_values)
198
+ hidden = outputs.last_hidden_state
199
+ tokens_to_skip = max(0, int(getattr(self.config, "vision_feature_prefix_tokens_to_skip", 1)))
200
+ if hidden.shape[1] > tokens_to_skip:
201
+ hidden = hidden[:, tokens_to_skip:, :]
202
+ return self.visual_projection(hidden)
203
+
204
+ def _build_layerwise_bias(
205
+ self,
206
+ anatomical_masks: Optional[torch.Tensor],
207
+ total_sequence_length: int,
208
+ vision_prefix_length: int,
209
+ ) -> Optional[torch.Tensor]:
210
+ if anatomical_masks is None:
211
+ return None
212
+ if getattr(self.config, "attention_bias_mode", "layerwise") == "gaussian_legacy":
213
+ vision_key_tokens = max(1, min(int(vision_prefix_length), int(total_sequence_length)))
214
+ legacy_bias = build_legacy_gaussian_attention_bias(
215
+ masks=anatomical_masks,
216
+ num_layers=self.config.num_attention_layers,
217
+ target_query_tokens=total_sequence_length,
218
+ target_key_tokens=vision_key_tokens,
219
+ base_kernel_size=self.config.layer_mask_base_kernel_size,
220
+ kernel_growth=self.config.layer_mask_kernel_growth,
221
+ strength=self.config.anatomical_attention_bias,
222
+ )
223
+ if vision_key_tokens == total_sequence_length:
224
+ return legacy_bias
225
+ full_bias = legacy_bias.new_zeros(
226
+ legacy_bias.shape[0],
227
+ legacy_bias.shape[1],
228
+ total_sequence_length,
229
+ total_sequence_length,
230
+ )
231
+ full_bias[:, :, :, :vision_key_tokens] = legacy_bias
232
+ return full_bias
233
+ return build_layerwise_attention_bias(
234
+ masks=anatomical_masks,
235
+ num_layers=self.config.num_attention_layers,
236
+ target_tokens=total_sequence_length,
237
+ base_kernel_size=self.config.layer_mask_base_kernel_size,
238
+ kernel_growth=self.config.layer_mask_kernel_growth,
239
+ strength=self.config.anatomical_attention_bias,
240
+ )
241
+
242
+ def _resolve_attention_bias(
243
+ self,
244
+ pixel_values: torch.Tensor,
245
+ anatomical_masks: Optional[torch.Tensor],
246
+ total_sequence_length: int,
247
+ vision_prefix_length: int,
248
+ ):
249
+ if anatomical_masks is not None:
250
+ return self._build_layerwise_bias(
251
+ anatomical_masks,
252
+ total_sequence_length=total_sequence_length,
253
+ vision_prefix_length=vision_prefix_length,
254
+ )
255
+ if self.segmenter is None:
256
+ return None
257
+ if getattr(self.config, "attention_bias_mode", "layerwise") == "gaussian_legacy":
258
+ combined_mask = self.segmenter.predict_mask(pixel_values)
259
+ if combined_mask is None:
260
+ logger.warning("Segmentation attention is enabled but no segmenter checkpoints were loaded; continuing without anatomical attention.")
261
+ return None
262
+ return self._build_layerwise_bias(
263
+ combined_mask,
264
+ total_sequence_length=total_sequence_length,
265
+ vision_prefix_length=vision_prefix_length,
266
+ )
267
+ layerwise_bias = self.segmenter(
268
+ pixel_values,
269
+ num_layers=self.config.num_attention_layers,
270
+ target_tokens=total_sequence_length,
271
+ strength=self.config.anatomical_attention_bias,
272
+ )
273
+ if layerwise_bias is None:
274
+ logger.warning("Segmentation attention is enabled but no segmenter checkpoints were loaded; continuing without anatomical attention.")
275
+ return layerwise_bias
276
+
277
+ def forward(
278
+ self,
279
+ pixel_values: torch.Tensor,
280
+ input_ids: Optional[torch.LongTensor] = None,
281
+ attention_mask: Optional[torch.Tensor] = None,
282
+ anatomical_masks: Optional[torch.Tensor] = None,
283
+ labels: Optional[torch.LongTensor] = None,
284
+ output_attentions: Optional[bool] = None,
285
+ output_hidden_states: Optional[bool] = None,
286
+ return_dict: Optional[bool] = True,
287
+ **kwargs,
288
+ ) -> LanaModelOutput:
289
+ vision_features = self._encode_images(pixel_values)
290
+ batch_size, prefix_length, _ = vision_features.shape
291
+
292
+ if input_ids is None:
293
+ bos = self.tokenizer.bos_token_id or self.tokenizer.eos_token_id
294
+ input_ids = torch.full((batch_size, 1), bos, device=vision_features.device, dtype=torch.long)
295
+ attention_mask = torch.ones_like(input_ids)
296
+ elif attention_mask is None:
297
+ attention_mask = torch.ones_like(input_ids)
298
+
299
+ text_embeds = self.text_decoder.transformer.wte(input_ids)
300
+ inputs_embeds = torch.cat([vision_features, text_embeds], dim=1)
301
+ merged_attention_mask = torch.cat(
302
+ [
303
+ torch.ones((batch_size, prefix_length), device=attention_mask.device, dtype=attention_mask.dtype),
304
+ attention_mask,
305
+ ],
306
+ dim=1,
307
+ )
308
+
309
+ merged_labels = None
310
+ if labels is not None:
311
+ ignore_prefix = torch.full((batch_size, prefix_length), -100, device=labels.device, dtype=labels.dtype)
312
+ merged_labels = torch.cat([ignore_prefix, labels], dim=1)
313
+
314
+ layerwise_bias = self._resolve_attention_bias(
315
+ pixel_values=pixel_values,
316
+ anatomical_masks=anatomical_masks,
317
+ total_sequence_length=inputs_embeds.shape[1],
318
+ vision_prefix_length=prefix_length,
319
+ )
320
+ decoder_outputs = self.text_decoder(
321
+ inputs_embeds=inputs_embeds,
322
+ attention_mask=merged_attention_mask,
323
+ labels=merged_labels,
324
+ segmentation_mask=layerwise_bias,
325
+ use_cache=False,
326
+ output_attentions=output_attentions,
327
+ output_hidden_states=output_hidden_states,
328
+ return_dict=True,
329
+ **kwargs,
330
+ )
331
+
332
+ return LanaModelOutput(
333
+ loss=decoder_outputs.loss,
334
+ logits=decoder_outputs.logits,
335
+ attentions=decoder_outputs.attentions,
336
+ layerwise_attentions=layerwise_bias,
337
+ hidden_states=decoder_outputs.hidden_states,
338
+ vision_features=vision_features,
339
+ )
340
+
341
+ @torch.inference_mode()
342
+ def generate(
343
+ self,
344
+ pixel_values: torch.Tensor,
345
+ anatomical_masks: Optional[torch.Tensor] = None,
346
+ max_new_tokens: int = 150,
347
+ **kwargs,
348
+ ):
349
+ vision_features = self._encode_images(pixel_values)
350
+ batch_size = pixel_values.shape[0]
351
+ if getattr(self.config, "generation_use_bos_token", True):
352
+ bos = self.tokenizer.bos_token_id or self.tokenizer.eos_token_id
353
+ start_tokens = torch.full((batch_size, 1), bos, device=pixel_values.device, dtype=torch.long)
354
+ text_embeds = self.text_decoder.transformer.wte(start_tokens)
355
+ inputs_embeds = torch.cat([vision_features, text_embeds], dim=1)
356
+ attention_mask = torch.ones(inputs_embeds.shape[:2], device=pixel_values.device, dtype=torch.long)
357
+ else:
358
+ inputs_embeds = vision_features
359
+ attention_mask = None
360
+
361
+ layerwise_bias = self._resolve_attention_bias(
362
+ pixel_values=pixel_values,
363
+ anatomical_masks=anatomical_masks,
364
+ total_sequence_length=inputs_embeds.shape[1] + max_new_tokens,
365
+ vision_prefix_length=vision_features.shape[1],
366
+ )
367
+ generation_kwargs = dict(
368
+ inputs_embeds=inputs_embeds,
369
+ max_new_tokens=max_new_tokens,
370
+ pad_token_id=self.tokenizer.pad_token_id,
371
+ do_sample=False,
372
+ num_beams=1,
373
+ segmentation_mask=layerwise_bias,
374
+ use_cache=True,
375
+ )
376
+ if attention_mask is not None:
377
+ generation_kwargs["attention_mask"] = attention_mask
378
+
379
+ repetition_penalty = float(getattr(self.config, "generation_repetition_penalty", 1.0))
380
+ if repetition_penalty > 1.0:
381
+ generation_kwargs["repetition_penalty"] = repetition_penalty
382
+
383
+ eos_token_id = self.tokenizer.eos_token_id
384
+ if getattr(self.config, "generation_stop_on_eos", False):
385
+ generation_kwargs["eos_token_id"] = eos_token_id
386
+ else:
387
+ generation_kwargs["eos_token_id"] = None
388
+ generation_kwargs["forced_eos_token_id"] = None
389
+ if eos_token_id is not None:
390
+ generation_kwargs["suppress_tokens"] = [int(eos_token_id)]
391
+
392
+ generation_kwargs.update(kwargs)
393
+ return self.text_decoder.generate(**generation_kwargs)
modeling_outputs.py ADDED
@@ -0,0 +1,15 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from dataclasses import dataclass
2
+ from typing import Optional, Tuple
3
+
4
+ import torch
5
+ from transformers.utils import ModelOutput
6
+
7
+
8
+ @dataclass
9
+ class LanaModelOutput(ModelOutput):
10
+ loss: Optional[torch.FloatTensor] = None
11
+ logits: Optional[torch.FloatTensor] = None
12
+ attentions: Optional[Tuple[torch.FloatTensor, ...]] = None
13
+ layerwise_attentions: Optional[torch.FloatTensor] = None
14
+ hidden_states: Optional[Tuple[torch.FloatTensor, ...]] = None
15
+ vision_features: Optional[torch.FloatTensor] = None
preprocessor_config.json ADDED
@@ -0,0 +1,27 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "do_convert_rgb": true,
3
+ "do_normalize": true,
4
+ "do_rescale": true,
5
+ "do_resize": true,
6
+ "image_mean": [
7
+ 0.485,
8
+ 0.456,
9
+ 0.406
10
+ ],
11
+ "image_processor_type": "LanaImageProcessor",
12
+ "image_std": [
13
+ 0.229,
14
+ 0.224,
15
+ 0.225
16
+ ],
17
+ "resample": 3,
18
+ "rescale_factor": 0.00392156862745098,
19
+ "size": {
20
+ "height": 512,
21
+ "width": 512
22
+ },
23
+ "auto_map": {
24
+ "AutoProcessor": "processing_lana.LanaProcessor"
25
+ },
26
+ "processor_class": "LanaProcessor"
27
+ }
processing_lana.py ADDED
@@ -0,0 +1,51 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ from pathlib import Path
4
+
5
+ from transformers import AutoTokenizer, GPT2Tokenizer
6
+ from transformers.processing_utils import ProcessorMixin
7
+
8
+ from .image_processing_lana import LanaImageProcessor
9
+
10
+
11
+ class LanaProcessor(ProcessorMixin):
12
+ attributes = ["image_processor", "tokenizer"]
13
+ image_processor_class = "LanaImageProcessor"
14
+ tokenizer_class = "AutoTokenizer"
15
+
16
+ def __init__(self, image_processor=None, tokenizer=None, **kwargs):
17
+ super().__init__(image_processor, tokenizer, **kwargs)
18
+
19
+ def __call__(self, images=None, text=None, **kwargs):
20
+ if images is None and text is None:
21
+ raise ValueError("LanaProcessor expected `images`, `text`, or both.")
22
+
23
+ encoded = {}
24
+ if images is not None:
25
+ encoded.update(self.image_processor(images=images, **kwargs))
26
+ if text is not None:
27
+ encoded.update(self.tokenizer(text, **kwargs))
28
+ return encoded
29
+
30
+ def batch_decode(self, *args, **kwargs):
31
+ return self.tokenizer.batch_decode(*args, **kwargs)
32
+
33
+ def decode(self, *args, **kwargs):
34
+ return self.tokenizer.decode(*args, **kwargs)
35
+
36
+ @classmethod
37
+ def from_pretrained(cls, pretrained_model_name_or_path, **kwargs):
38
+ kwargs = dict(kwargs)
39
+ kwargs.pop("trust_remote_code", None)
40
+ image_processor = LanaImageProcessor.from_pretrained(pretrained_model_name_or_path, **kwargs)
41
+ source = Path(str(pretrained_model_name_or_path))
42
+ if source.exists():
43
+ tokenizer = GPT2Tokenizer.from_pretrained(pretrained_model_name_or_path)
44
+ else:
45
+ tokenizer = AutoTokenizer.from_pretrained(
46
+ pretrained_model_name_or_path,
47
+ trust_remote_code=True,
48
+ use_fast=False,
49
+ **kwargs,
50
+ )
51
+ return cls(image_processor=image_processor, tokenizer=tokenizer)
processor_config.json ADDED
@@ -0,0 +1,29 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "image_processor": {
3
+ "do_resize": true,
4
+ "size": {
5
+ "height": 512,
6
+ "width": 512
7
+ },
8
+ "resample": 3,
9
+ "do_rescale": true,
10
+ "rescale_factor": 0.00392156862745098,
11
+ "do_normalize": true,
12
+ "image_mean": [
13
+ 0.485,
14
+ 0.456,
15
+ 0.406
16
+ ],
17
+ "image_std": [
18
+ 0.229,
19
+ 0.224,
20
+ 0.225
21
+ ],
22
+ "do_convert_rgb": true,
23
+ "image_processor_type": "LanaImageProcessor"
24
+ },
25
+ "processor_class": "LanaProcessor",
26
+ "auto_map": {
27
+ "AutoProcessor": "processing_lana.LanaProcessor"
28
+ }
29
+ }
segmenters.py ADDED
@@ -0,0 +1,147 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import logging
2
+ from pathlib import Path
3
+
4
+ import torch
5
+ import torch.nn as nn
6
+ from transformers import AutoConfig, AutoModel
7
+
8
+ from .layerwise_anatomical_attention import build_layerwise_attention_bias
9
+
10
+ LOGGER = logging.getLogger(__name__)
11
+
12
+
13
+ def _freeze_module(module: nn.Module) -> None:
14
+ for param in module.parameters():
15
+ param.requires_grad = False
16
+
17
+
18
+ class _DinoUNetLung(nn.Module):
19
+ def __init__(self, model_name: str, freeze: bool = True, load_pretrained: bool = True):
20
+ super().__init__()
21
+ if load_pretrained:
22
+ self.encoder = AutoModel.from_pretrained(model_name, trust_remote_code=True)
23
+ else:
24
+ self.encoder = AutoModel.from_config(AutoConfig.from_pretrained(model_name, trust_remote_code=True), trust_remote_code=True)
25
+ self.channel_adapter = nn.Conv2d(768, 512, kernel_size=1)
26
+ self.decoder = nn.Sequential(
27
+ nn.Conv2d(512, 256, 3, padding=1),
28
+ nn.ReLU(inplace=True),
29
+ nn.ConvTranspose2d(256, 128, 2, stride=2),
30
+ nn.ReLU(inplace=True),
31
+ nn.ConvTranspose2d(128, 64, 2, stride=2),
32
+ nn.ReLU(inplace=True),
33
+ nn.Conv2d(64, 1, 1),
34
+ )
35
+ if freeze:
36
+ _freeze_module(self)
37
+
38
+ @torch.no_grad()
39
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
40
+ enc_feats = self.encoder(x, output_hidden_states=True, return_dict=True)
41
+ feats = next(h for h in reversed(enc_feats.hidden_states) if isinstance(h, torch.Tensor) and h.ndim == 4)
42
+ feats = self.channel_adapter(feats)
43
+ pred = self.decoder(feats)
44
+ return (torch.sigmoid(pred) > 0.5).float()
45
+
46
+
47
+ class _DinoUNetHeart(nn.Module):
48
+ def __init__(self, model_name: str, freeze: bool = True, load_pretrained: bool = True):
49
+ super().__init__()
50
+ if load_pretrained:
51
+ self.encoder = AutoModel.from_pretrained(model_name, trust_remote_code=True)
52
+ else:
53
+ self.encoder = AutoModel.from_config(AutoConfig.from_pretrained(model_name, trust_remote_code=True), trust_remote_code=True)
54
+ self.adapter = nn.Conv2d(768, 512, 1)
55
+ self.decoder = nn.Sequential(
56
+ nn.Conv2d(512, 256, 3, padding=1),
57
+ nn.ReLU(True),
58
+ nn.ConvTranspose2d(256, 128, 2, 2),
59
+ nn.ReLU(True),
60
+ nn.ConvTranspose2d(128, 64, 2, 2),
61
+ nn.ReLU(True),
62
+ nn.Conv2d(64, 3, 1),
63
+ )
64
+ if freeze:
65
+ _freeze_module(self)
66
+
67
+ @torch.no_grad()
68
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
69
+ enc = self.encoder(x, output_hidden_states=True, return_dict=True)
70
+ feat = next(h for h in reversed(enc.hidden_states) if isinstance(h, torch.Tensor) and h.ndim == 4)
71
+ feat = self.adapter(feat)
72
+ logits = self.decoder(feat)
73
+ pred = torch.argmax(logits, dim=1)
74
+ return (pred == 2).unsqueeze(1).float()
75
+
76
+
77
+ class AnatomicalSegmenter(nn.Module):
78
+ def __init__(
79
+ self,
80
+ model_name: str,
81
+ freeze: bool = True,
82
+ lung_checkpoint: str = "",
83
+ heart_checkpoint: str = "",
84
+ load_pretrained: bool = True,
85
+ assume_weights_from_model_state: bool = False,
86
+ ):
87
+ super().__init__()
88
+ self.lung_model = _DinoUNetLung(model_name=model_name, freeze=freeze, load_pretrained=load_pretrained)
89
+ self.heart_model = _DinoUNetHeart(model_name=model_name, freeze=freeze, load_pretrained=load_pretrained)
90
+ if assume_weights_from_model_state:
91
+ self.loaded_lung_checkpoint = True
92
+ self.loaded_heart_checkpoint = True
93
+ else:
94
+ self.loaded_lung_checkpoint = self._load_submodule(self.lung_model, lung_checkpoint, "lung")
95
+ self.loaded_heart_checkpoint = self._load_submodule(self.heart_model, heart_checkpoint, "heart")
96
+
97
+ @staticmethod
98
+ def _load_submodule(module: nn.Module, checkpoint_path: str, label: str) -> bool:
99
+ if not checkpoint_path:
100
+ return False
101
+ path = Path(checkpoint_path)
102
+ if not path.exists():
103
+ LOGGER.warning("Requested %s segmenter checkpoint does not exist: %s", label, path)
104
+ return False
105
+ if any(getattr(param, "is_meta", False) for param in module.parameters()):
106
+ LOGGER.info(
107
+ "Deferring %s segmenter checkpoint preload for meta-initialized module; packaged model weights will finish loading it.",
108
+ label,
109
+ )
110
+ return True
111
+ state = torch.load(path, map_location="cpu", weights_only=False)
112
+ if isinstance(state, dict) and "state_dict" in state:
113
+ state = state["state_dict"]
114
+ module.load_state_dict(state, strict=False)
115
+ LOGGER.info("Loaded %s segmenter checkpoint from %s", label, path)
116
+ return True
117
+
118
+ @property
119
+ def has_any_checkpoint(self) -> bool:
120
+ return self.loaded_lung_checkpoint or self.loaded_heart_checkpoint
121
+
122
+ @torch.no_grad()
123
+ def predict_mask(self, pixel_values: torch.Tensor) -> torch.Tensor | None:
124
+ if not self.has_any_checkpoint:
125
+ return None
126
+
127
+ masks = []
128
+ if self.loaded_heart_checkpoint:
129
+ masks.append(self.heart_model(pixel_values))
130
+ if self.loaded_lung_checkpoint:
131
+ masks.append(self.lung_model(pixel_values))
132
+ if not masks:
133
+ return None
134
+
135
+ return torch.clamp(sum(masks), 0.0, 1.0)
136
+
137
+ @torch.no_grad()
138
+ def forward(self, pixel_values: torch.Tensor, num_layers: int, target_tokens: int, strength: float) -> torch.Tensor | None:
139
+ combined_mask = self.predict_mask(pixel_values)
140
+ if combined_mask is None:
141
+ return None
142
+ return build_layerwise_attention_bias(
143
+ masks=combined_mask,
144
+ num_layers=num_layers,
145
+ target_tokens=target_tokens,
146
+ strength=strength,
147
+ )
tokenizer.json ADDED
The diff for this file is too large to render. See raw diff
 
tokenizer_config.json ADDED
@@ -0,0 +1,17 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "add_prefix_space": false,
3
+ "backend": "tokenizers",
4
+ "bos_token": "<|endoftext|>",
5
+ "eos_token": "<|endoftext|>",
6
+ "errors": "replace",
7
+ "is_local": true,
8
+ "max_length": 1024,
9
+ "model_max_length": 1024,
10
+ "pad_token": "<|endoftext|>",
11
+ "processor_class": "LanaProcessor",
12
+ "stride": 0,
13
+ "tokenizer_class": "GPT2Tokenizer",
14
+ "truncation_side": "right",
15
+ "truncation_strategy": "longest_first",
16
+ "unk_token": "<|endoftext|>"
17
+ }
vocab.json ADDED
The diff for this file is too large to render. See raw diff