Image-Text-to-Video
Diffusers
Safetensors
MiniMaxH3ModularPipeline
text-to-video
image-to-video
video-to-video
text-to-audio-video
image-to-audio-video
image-text-to-audio-video
video-to-audio-video
audio-to-audio-video
audio-video-generation
multimodal
synchronized-audio-video
reference-to-audio-video
Instructions to use unsloth/MiniMax-H3 with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Diffusers
How to use unsloth/MiniMax-H3 with Diffusers:
pip install -U diffusers transformers accelerate
import torch from diffusers import DiffusionPipeline # switch to "mps" for apple devices pipe = DiffusionPipeline.from_pretrained("unsloth/MiniMax-H3", dtype=torch.bfloat16, device_map="cuda") prompt = "Astronaut in a jungle, cold color palette, muted colors, detailed, 8k" image = pipe(prompt).images[0] - Notebooks
- Google Colab
- Kaggle
Add files using upload-large-folder tool
Browse filesThis view is limited to 50 files because it contains too many changes. See raw diff
- FL2VA/processor/video_preprocessor_config.json +21 -0
- FL2VA/processor/vocab.json +0 -0
- FL2VA/text_encoder/chat_template.json +3 -0
- FL2VA/text_encoder/config.json +62 -0
- FL2VA/text_encoder/merges.txt +0 -0
- FL2VA/text_encoder/model.safetensors.index.json +1065 -0
- FL2VA/text_encoder/preprocessor_config.json +21 -0
- FL2VA/text_encoder/tokenizer.json +0 -0
- FL2VA/text_encoder/tokenizer_config.json +246 -0
- FL2VA/text_encoder/video_preprocessor_config.json +21 -0
- FL2VA/text_encoder/vocab.json +0 -0
- FL2VA/tokenizer/merges.txt +0 -0
- FL2VA/tokenizer/tokenizer.json +0 -0
- FL2VA/tokenizer/tokenizer_config.json +246 -0
- FL2VA/tokenizer/vocab.json +0 -0
- FL2VA/transformer/config.json +27 -0
- FL2VA/transformer/model-00013-of-00013.safetensors +3 -0
- FL2VA/transformer/model.safetensors.index.json +542 -0
- FL2VA/video_vae/attention.py +163 -0
- FL2VA/video_vae/base_module.py +282 -0
- FL2VA/video_vae/config.json +74 -0
- FL2VA/video_vae/conv.py +159 -0
- FL2VA/video_vae/flash.py +178 -0
- FL2VA/video_vae/func.py +163 -0
- FL2VA/video_vae/klvae.py +1258 -0
- FL2VA/video_vae/minimax_h3_video_vae.py +122 -0
- FL2VA/video_vae/norm.py +357 -0
- FL2VA/video_vae/normalize.py +39 -0
- FL2VA/video_vae/parallel.py +418 -0
- FL2VA/video_vae/source/config.json +71 -0
- FL2VA/video_vae/source/model.safetensors +3 -0
- FL2VA/video_vae/utils.py +18 -0
- FL2VA/video_vae/vae_cnn.py +304 -0
- FL2VA/video_vae/vae_module.py +53 -0
- FL2VA/video_vae/vae_processor.py +234 -0
- FL2VA/video_vae/vae_vit.py +380 -0
- Ref2VA/audio_vae/model.safetensors +3 -0
- assets/fl2va.mp4 +3 -0
- assets/full-arch.png +3 -0
- assets/h3_direct_2k.mp4 +3 -0
- assets/h3_direct_768p.mp4 +3 -0
- assets/i2va.mp4 +3 -0
- assets/i2va_2k.mp4 +3 -0
- assets/i2va_direct_2k.mp4 +3 -0
- assets/i2va_direct_768p.mp4 +3 -0
- assets/minimax-h3.png +3 -0
- assets/overview.png +3 -0
- assets/r2va.mp4 +3 -0
- assets/r2va_2k.mp4 +3 -0
- assets/r2va_direct_2k.mp4 +3 -0
FL2VA/processor/video_preprocessor_config.json
ADDED
|
@@ -0,0 +1,21 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"size": {
|
| 3 |
+
"longest_edge": 25165824,
|
| 4 |
+
"shortest_edge": 4096
|
| 5 |
+
},
|
| 6 |
+
"patch_size": 16,
|
| 7 |
+
"temporal_patch_size": 2,
|
| 8 |
+
"merge_size": 2,
|
| 9 |
+
"image_mean": [
|
| 10 |
+
0.5,
|
| 11 |
+
0.5,
|
| 12 |
+
0.5
|
| 13 |
+
],
|
| 14 |
+
"image_std": [
|
| 15 |
+
0.5,
|
| 16 |
+
0.5,
|
| 17 |
+
0.5
|
| 18 |
+
],
|
| 19 |
+
"processor_class": "Qwen3VLProcessor",
|
| 20 |
+
"video_processor_type": "Qwen3VLVideoProcessor"
|
| 21 |
+
}
|
FL2VA/processor/vocab.json
ADDED
|
The diff for this file is too large to render.
See raw diff
|
|
|
FL2VA/text_encoder/chat_template.json
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"chat_template": "{%- if tools %}\n {{- '<|im_start|>system\\n' }}\n {%- if messages[0].role == 'system' %}\n {%- if messages[0].content is string %}\n {{- messages[0].content }}\n {%- else %}\n {%- for content in messages[0].content %}\n {%- if 'text' in content %}\n {{- content.text }}\n {%- endif %}\n {%- endfor %}\n {%- endif %}\n {{- '\\n\\n' }}\n {%- endif %}\n {{- \"# Tools\\n\\nYou may call one or more functions to assist with the user query.\\n\\nYou are provided with function signatures within <tools></tools> XML tags:\\n<tools>\" }}\n {%- for tool in tools %}\n {{- \"\\n\" }}\n {{- tool | tojson }}\n {%- endfor %}\n {{- \"\\n</tools>\\n\\nFor each function call, return a json object with function name and arguments within <tool_call></tool_call> XML tags:\\n<tool_call>\\n{\\\"name\\\": <function-name>, \\\"arguments\\\": <args-json-object>}\\n</tool_call><|im_end|>\\n\" }}\n{%- else %}\n {%- if messages[0].role == 'system' %}\n {{- '<|im_start|>system\\n' }}\n {%- if messages[0].content is string %}\n {{- messages[0].content }}\n {%- else %}\n {%- for content in messages[0].content %}\n {%- if 'text' in content %}\n {{- content.text }}\n {%- endif %}\n {%- endfor %}\n {%- endif %}\n {{- '<|im_end|>\\n' }}\n {%- endif %}\n{%- endif %}\n{%- set image_count = namespace(value=0) %}\n{%- set video_count = namespace(value=0) %}\n{%- for message in messages %}\n {%- if message.role == \"user\" %}\n {{- '<|im_start|>' + message.role + '\\n' }}\n {%- if message.content is string %}\n {{- message.content }}\n {%- else %}\n {%- for content in message.content %}\n {%- if content.type == 'image' or 'image' in content or 'image_url' in content %}\n {%- set image_count.value = image_count.value + 1 %}\n {%- if add_vision_id %}Picture {{ image_count.value }}: {% endif -%}\n <|vision_start|><|image_pad|><|vision_end|>\n {%- elif content.type == 'video' or 'video' in content %}\n {%- set video_count.value = video_count.value + 1 %}\n {%- if add_vision_id %}Video {{ video_count.value }}: {% endif -%}\n <|vision_start|><|video_pad|><|vision_end|>\n {%- elif 'text' in content %}\n {{- content.text }}\n {%- endif %}\n {%- endfor %}\n {%- endif %}\n {{- '<|im_end|>\\n' }}\n {%- elif message.role == \"assistant\" %}\n {{- '<|im_start|>' + message.role + '\\n' }}\n {%- if message.content is string %}\n {{- message.content }}\n {%- else %}\n {%- for content_item in message.content %}\n {%- if 'text' in content_item %}\n {{- content_item.text }}\n {%- endif %}\n {%- endfor %}\n {%- endif %}\n {%- if message.tool_calls %}\n {%- for tool_call in message.tool_calls %}\n {%- if (loop.first and message.content) or (not loop.first) %}\n {{- '\\n' }}\n {%- endif %}\n {%- if tool_call.function %}\n {%- set tool_call = tool_call.function %}\n {%- endif %}\n {{- '<tool_call>\\n{\"name\": \"' }}\n {{- tool_call.name }}\n {{- '\", \"arguments\": ' }}\n {%- if tool_call.arguments is string %}\n {{- tool_call.arguments }}\n {%- else %}\n {{- tool_call.arguments | tojson }}\n {%- endif %}\n {{- '}\\n</tool_call>' }}\n {%- endfor %}\n {%- endif %}\n {{- '<|im_end|>\\n' }}\n {%- elif message.role == \"tool\" %}\n {%- if loop.first or (messages[loop.index0 - 1].role != \"tool\") %}\n {{- '<|im_start|>user' }}\n {%- endif %}\n {{- '\\n<tool_response>\\n' }}\n {%- if message.content is string %}\n {{- message.content }}\n {%- else %}\n {%- for content in message.content %}\n {%- if content.type == 'image' or 'image' in content or 'image_url' in content %}\n {%- set image_count.value = image_count.value + 1 %}\n {%- if add_vision_id %}Picture {{ image_count.value }}: {% endif -%}\n <|vision_start|><|image_pad|><|vision_end|>\n {%- elif content.type == 'video' or 'video' in content %}\n {%- set video_count.value = video_count.value + 1 %}\n {%- if add_vision_id %}Video {{ video_count.value }}: {% endif -%}\n <|vision_start|><|video_pad|><|vision_end|>\n {%- elif 'text' in content %}\n {{- content.text }}\n {%- endif %}\n {%- endfor %}\n {%- endif %}\n {{- '\\n</tool_response>' }}\n {%- if loop.last or (messages[loop.index0 + 1].role != \"tool\") %}\n {{- '<|im_end|>\\n' }}\n {%- endif %}\n {%- endif %}\n{%- endfor %}\n{%- if add_generation_prompt %}\n {{- '<|im_start|>assistant\\n' }}\n{%- endif %}\n"
|
| 3 |
+
}
|
FL2VA/text_encoder/config.json
ADDED
|
@@ -0,0 +1,62 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"architectures": [
|
| 3 |
+
"Qwen3VLForConditionalGeneration"
|
| 4 |
+
],
|
| 5 |
+
"image_token_id": 151655,
|
| 6 |
+
"model_type": "qwen3_vl",
|
| 7 |
+
"text_config": {
|
| 8 |
+
"attention_bias": false,
|
| 9 |
+
"attention_dropout": 0.0,
|
| 10 |
+
"bos_token_id": 151643,
|
| 11 |
+
"dtype": "bfloat16",
|
| 12 |
+
"eos_token_id": 151645,
|
| 13 |
+
"head_dim": 128,
|
| 14 |
+
"hidden_act": "silu",
|
| 15 |
+
"hidden_size": 5120,
|
| 16 |
+
"initializer_range": 0.02,
|
| 17 |
+
"intermediate_size": 25600,
|
| 18 |
+
"max_position_embeddings": 262144,
|
| 19 |
+
"model_type": "qwen3_vl_text",
|
| 20 |
+
"num_attention_heads": 64,
|
| 21 |
+
"num_hidden_layers": 64,
|
| 22 |
+
"num_key_value_heads": 8,
|
| 23 |
+
"rms_norm_eps": 1e-06,
|
| 24 |
+
"rope_scaling": {
|
| 25 |
+
"mrope_interleaved": true,
|
| 26 |
+
"mrope_section": [
|
| 27 |
+
24,
|
| 28 |
+
20,
|
| 29 |
+
20
|
| 30 |
+
],
|
| 31 |
+
"rope_type": "default"
|
| 32 |
+
},
|
| 33 |
+
"rope_theta": 5000000,
|
| 34 |
+
"use_cache": true,
|
| 35 |
+
"vocab_size": 151936
|
| 36 |
+
},
|
| 37 |
+
"tie_word_embeddings": false,
|
| 38 |
+
"transformers_version": "4.57.0.dev0",
|
| 39 |
+
"video_token_id": 151656,
|
| 40 |
+
"vision_config": {
|
| 41 |
+
"deepstack_visual_indexes": [
|
| 42 |
+
8,
|
| 43 |
+
16,
|
| 44 |
+
24
|
| 45 |
+
],
|
| 46 |
+
"depth": 27,
|
| 47 |
+
"hidden_act": "gelu_pytorch_tanh",
|
| 48 |
+
"hidden_size": 1152,
|
| 49 |
+
"in_channels": 3,
|
| 50 |
+
"initializer_range": 0.02,
|
| 51 |
+
"intermediate_size": 4304,
|
| 52 |
+
"model_type": "qwen3_vl",
|
| 53 |
+
"num_heads": 16,
|
| 54 |
+
"num_position_embeddings": 2304,
|
| 55 |
+
"out_hidden_size": 5120,
|
| 56 |
+
"patch_size": 16,
|
| 57 |
+
"spatial_merge_size": 2,
|
| 58 |
+
"temporal_patch_size": 2
|
| 59 |
+
},
|
| 60 |
+
"vision_end_token_id": 151653,
|
| 61 |
+
"vision_start_token_id": 151652
|
| 62 |
+
}
|
FL2VA/text_encoder/merges.txt
ADDED
|
The diff for this file is too large to render.
See raw diff
|
|
|
FL2VA/text_encoder/model.safetensors.index.json
ADDED
|
@@ -0,0 +1,1065 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"metadata": {
|
| 3 |
+
"total_size": 66714780128
|
| 4 |
+
},
|
| 5 |
+
"weight_map": {
|
| 6 |
+
"lm_head.weight": "model-00014-of-00014.safetensors",
|
| 7 |
+
"model.language_model.embed_tokens.weight": "model-00001-of-00014.safetensors",
|
| 8 |
+
"model.language_model.layers.0.input_layernorm.weight": "model-00001-of-00014.safetensors",
|
| 9 |
+
"model.language_model.layers.0.mlp.down_proj.weight": "model-00001-of-00014.safetensors",
|
| 10 |
+
"model.language_model.layers.0.mlp.gate_proj.weight": "model-00001-of-00014.safetensors",
|
| 11 |
+
"model.language_model.layers.0.mlp.up_proj.weight": "model-00001-of-00014.safetensors",
|
| 12 |
+
"model.language_model.layers.0.post_attention_layernorm.weight": "model-00001-of-00014.safetensors",
|
| 13 |
+
"model.language_model.layers.0.self_attn.k_norm.weight": "model-00001-of-00014.safetensors",
|
| 14 |
+
"model.language_model.layers.0.self_attn.k_proj.weight": "model-00001-of-00014.safetensors",
|
| 15 |
+
"model.language_model.layers.0.self_attn.o_proj.weight": "model-00001-of-00014.safetensors",
|
| 16 |
+
"model.language_model.layers.0.self_attn.q_norm.weight": "model-00001-of-00014.safetensors",
|
| 17 |
+
"model.language_model.layers.0.self_attn.q_proj.weight": "model-00001-of-00014.safetensors",
|
| 18 |
+
"model.language_model.layers.0.self_attn.v_proj.weight": "model-00001-of-00014.safetensors",
|
| 19 |
+
"model.language_model.layers.1.input_layernorm.weight": "model-00001-of-00014.safetensors",
|
| 20 |
+
"model.language_model.layers.1.mlp.down_proj.weight": "model-00001-of-00014.safetensors",
|
| 21 |
+
"model.language_model.layers.1.mlp.gate_proj.weight": "model-00001-of-00014.safetensors",
|
| 22 |
+
"model.language_model.layers.1.mlp.up_proj.weight": "model-00001-of-00014.safetensors",
|
| 23 |
+
"model.language_model.layers.1.post_attention_layernorm.weight": "model-00001-of-00014.safetensors",
|
| 24 |
+
"model.language_model.layers.1.self_attn.k_norm.weight": "model-00001-of-00014.safetensors",
|
| 25 |
+
"model.language_model.layers.1.self_attn.k_proj.weight": "model-00001-of-00014.safetensors",
|
| 26 |
+
"model.language_model.layers.1.self_attn.o_proj.weight": "model-00001-of-00014.safetensors",
|
| 27 |
+
"model.language_model.layers.1.self_attn.q_norm.weight": "model-00001-of-00014.safetensors",
|
| 28 |
+
"model.language_model.layers.1.self_attn.q_proj.weight": "model-00001-of-00014.safetensors",
|
| 29 |
+
"model.language_model.layers.1.self_attn.v_proj.weight": "model-00001-of-00014.safetensors",
|
| 30 |
+
"model.language_model.layers.10.input_layernorm.weight": "model-00003-of-00014.safetensors",
|
| 31 |
+
"model.language_model.layers.10.mlp.down_proj.weight": "model-00003-of-00014.safetensors",
|
| 32 |
+
"model.language_model.layers.10.mlp.gate_proj.weight": "model-00003-of-00014.safetensors",
|
| 33 |
+
"model.language_model.layers.10.mlp.up_proj.weight": "model-00003-of-00014.safetensors",
|
| 34 |
+
"model.language_model.layers.10.post_attention_layernorm.weight": "model-00003-of-00014.safetensors",
|
| 35 |
+
"model.language_model.layers.10.self_attn.k_norm.weight": "model-00003-of-00014.safetensors",
|
| 36 |
+
"model.language_model.layers.10.self_attn.k_proj.weight": "model-00003-of-00014.safetensors",
|
| 37 |
+
"model.language_model.layers.10.self_attn.o_proj.weight": "model-00003-of-00014.safetensors",
|
| 38 |
+
"model.language_model.layers.10.self_attn.q_norm.weight": "model-00003-of-00014.safetensors",
|
| 39 |
+
"model.language_model.layers.10.self_attn.q_proj.weight": "model-00003-of-00014.safetensors",
|
| 40 |
+
"model.language_model.layers.10.self_attn.v_proj.weight": "model-00003-of-00014.safetensors",
|
| 41 |
+
"model.language_model.layers.11.input_layernorm.weight": "model-00003-of-00014.safetensors",
|
| 42 |
+
"model.language_model.layers.11.mlp.down_proj.weight": "model-00003-of-00014.safetensors",
|
| 43 |
+
"model.language_model.layers.11.mlp.gate_proj.weight": "model-00003-of-00014.safetensors",
|
| 44 |
+
"model.language_model.layers.11.mlp.up_proj.weight": "model-00003-of-00014.safetensors",
|
| 45 |
+
"model.language_model.layers.11.post_attention_layernorm.weight": "model-00003-of-00014.safetensors",
|
| 46 |
+
"model.language_model.layers.11.self_attn.k_norm.weight": "model-00003-of-00014.safetensors",
|
| 47 |
+
"model.language_model.layers.11.self_attn.k_proj.weight": "model-00003-of-00014.safetensors",
|
| 48 |
+
"model.language_model.layers.11.self_attn.o_proj.weight": "model-00003-of-00014.safetensors",
|
| 49 |
+
"model.language_model.layers.11.self_attn.q_norm.weight": "model-00003-of-00014.safetensors",
|
| 50 |
+
"model.language_model.layers.11.self_attn.q_proj.weight": "model-00003-of-00014.safetensors",
|
| 51 |
+
"model.language_model.layers.11.self_attn.v_proj.weight": "model-00003-of-00014.safetensors",
|
| 52 |
+
"model.language_model.layers.12.input_layernorm.weight": "model-00003-of-00014.safetensors",
|
| 53 |
+
"model.language_model.layers.12.mlp.down_proj.weight": "model-00003-of-00014.safetensors",
|
| 54 |
+
"model.language_model.layers.12.mlp.gate_proj.weight": "model-00003-of-00014.safetensors",
|
| 55 |
+
"model.language_model.layers.12.mlp.up_proj.weight": "model-00003-of-00014.safetensors",
|
| 56 |
+
"model.language_model.layers.12.post_attention_layernorm.weight": "model-00003-of-00014.safetensors",
|
| 57 |
+
"model.language_model.layers.12.self_attn.k_norm.weight": "model-00003-of-00014.safetensors",
|
| 58 |
+
"model.language_model.layers.12.self_attn.k_proj.weight": "model-00003-of-00014.safetensors",
|
| 59 |
+
"model.language_model.layers.12.self_attn.o_proj.weight": "model-00003-of-00014.safetensors",
|
| 60 |
+
"model.language_model.layers.12.self_attn.q_norm.weight": "model-00003-of-00014.safetensors",
|
| 61 |
+
"model.language_model.layers.12.self_attn.q_proj.weight": "model-00003-of-00014.safetensors",
|
| 62 |
+
"model.language_model.layers.12.self_attn.v_proj.weight": "model-00003-of-00014.safetensors",
|
| 63 |
+
"model.language_model.layers.13.input_layernorm.weight": "model-00003-of-00014.safetensors",
|
| 64 |
+
"model.language_model.layers.13.mlp.down_proj.weight": "model-00004-of-00014.safetensors",
|
| 65 |
+
"model.language_model.layers.13.mlp.gate_proj.weight": "model-00003-of-00014.safetensors",
|
| 66 |
+
"model.language_model.layers.13.mlp.up_proj.weight": "model-00004-of-00014.safetensors",
|
| 67 |
+
"model.language_model.layers.13.post_attention_layernorm.weight": "model-00003-of-00014.safetensors",
|
| 68 |
+
"model.language_model.layers.13.self_attn.k_norm.weight": "model-00003-of-00014.safetensors",
|
| 69 |
+
"model.language_model.layers.13.self_attn.k_proj.weight": "model-00003-of-00014.safetensors",
|
| 70 |
+
"model.language_model.layers.13.self_attn.o_proj.weight": "model-00003-of-00014.safetensors",
|
| 71 |
+
"model.language_model.layers.13.self_attn.q_norm.weight": "model-00003-of-00014.safetensors",
|
| 72 |
+
"model.language_model.layers.13.self_attn.q_proj.weight": "model-00003-of-00014.safetensors",
|
| 73 |
+
"model.language_model.layers.13.self_attn.v_proj.weight": "model-00003-of-00014.safetensors",
|
| 74 |
+
"model.language_model.layers.14.input_layernorm.weight": "model-00004-of-00014.safetensors",
|
| 75 |
+
"model.language_model.layers.14.mlp.down_proj.weight": "model-00004-of-00014.safetensors",
|
| 76 |
+
"model.language_model.layers.14.mlp.gate_proj.weight": "model-00004-of-00014.safetensors",
|
| 77 |
+
"model.language_model.layers.14.mlp.up_proj.weight": "model-00004-of-00014.safetensors",
|
| 78 |
+
"model.language_model.layers.14.post_attention_layernorm.weight": "model-00004-of-00014.safetensors",
|
| 79 |
+
"model.language_model.layers.14.self_attn.k_norm.weight": "model-00004-of-00014.safetensors",
|
| 80 |
+
"model.language_model.layers.14.self_attn.k_proj.weight": "model-00004-of-00014.safetensors",
|
| 81 |
+
"model.language_model.layers.14.self_attn.o_proj.weight": "model-00004-of-00014.safetensors",
|
| 82 |
+
"model.language_model.layers.14.self_attn.q_norm.weight": "model-00004-of-00014.safetensors",
|
| 83 |
+
"model.language_model.layers.14.self_attn.q_proj.weight": "model-00004-of-00014.safetensors",
|
| 84 |
+
"model.language_model.layers.14.self_attn.v_proj.weight": "model-00004-of-00014.safetensors",
|
| 85 |
+
"model.language_model.layers.15.input_layernorm.weight": "model-00004-of-00014.safetensors",
|
| 86 |
+
"model.language_model.layers.15.mlp.down_proj.weight": "model-00004-of-00014.safetensors",
|
| 87 |
+
"model.language_model.layers.15.mlp.gate_proj.weight": "model-00004-of-00014.safetensors",
|
| 88 |
+
"model.language_model.layers.15.mlp.up_proj.weight": "model-00004-of-00014.safetensors",
|
| 89 |
+
"model.language_model.layers.15.post_attention_layernorm.weight": "model-00004-of-00014.safetensors",
|
| 90 |
+
"model.language_model.layers.15.self_attn.k_norm.weight": "model-00004-of-00014.safetensors",
|
| 91 |
+
"model.language_model.layers.15.self_attn.k_proj.weight": "model-00004-of-00014.safetensors",
|
| 92 |
+
"model.language_model.layers.15.self_attn.o_proj.weight": "model-00004-of-00014.safetensors",
|
| 93 |
+
"model.language_model.layers.15.self_attn.q_norm.weight": "model-00004-of-00014.safetensors",
|
| 94 |
+
"model.language_model.layers.15.self_attn.q_proj.weight": "model-00004-of-00014.safetensors",
|
| 95 |
+
"model.language_model.layers.15.self_attn.v_proj.weight": "model-00004-of-00014.safetensors",
|
| 96 |
+
"model.language_model.layers.16.input_layernorm.weight": "model-00004-of-00014.safetensors",
|
| 97 |
+
"model.language_model.layers.16.mlp.down_proj.weight": "model-00004-of-00014.safetensors",
|
| 98 |
+
"model.language_model.layers.16.mlp.gate_proj.weight": "model-00004-of-00014.safetensors",
|
| 99 |
+
"model.language_model.layers.16.mlp.up_proj.weight": "model-00004-of-00014.safetensors",
|
| 100 |
+
"model.language_model.layers.16.post_attention_layernorm.weight": "model-00004-of-00014.safetensors",
|
| 101 |
+
"model.language_model.layers.16.self_attn.k_norm.weight": "model-00004-of-00014.safetensors",
|
| 102 |
+
"model.language_model.layers.16.self_attn.k_proj.weight": "model-00004-of-00014.safetensors",
|
| 103 |
+
"model.language_model.layers.16.self_attn.o_proj.weight": "model-00004-of-00014.safetensors",
|
| 104 |
+
"model.language_model.layers.16.self_attn.q_norm.weight": "model-00004-of-00014.safetensors",
|
| 105 |
+
"model.language_model.layers.16.self_attn.q_proj.weight": "model-00004-of-00014.safetensors",
|
| 106 |
+
"model.language_model.layers.16.self_attn.v_proj.weight": "model-00004-of-00014.safetensors",
|
| 107 |
+
"model.language_model.layers.17.input_layernorm.weight": "model-00004-of-00014.safetensors",
|
| 108 |
+
"model.language_model.layers.17.mlp.down_proj.weight": "model-00004-of-00014.safetensors",
|
| 109 |
+
"model.language_model.layers.17.mlp.gate_proj.weight": "model-00004-of-00014.safetensors",
|
| 110 |
+
"model.language_model.layers.17.mlp.up_proj.weight": "model-00004-of-00014.safetensors",
|
| 111 |
+
"model.language_model.layers.17.post_attention_layernorm.weight": "model-00004-of-00014.safetensors",
|
| 112 |
+
"model.language_model.layers.17.self_attn.k_norm.weight": "model-00004-of-00014.safetensors",
|
| 113 |
+
"model.language_model.layers.17.self_attn.k_proj.weight": "model-00004-of-00014.safetensors",
|
| 114 |
+
"model.language_model.layers.17.self_attn.o_proj.weight": "model-00004-of-00014.safetensors",
|
| 115 |
+
"model.language_model.layers.17.self_attn.q_norm.weight": "model-00004-of-00014.safetensors",
|
| 116 |
+
"model.language_model.layers.17.self_attn.q_proj.weight": "model-00004-of-00014.safetensors",
|
| 117 |
+
"model.language_model.layers.17.self_attn.v_proj.weight": "model-00004-of-00014.safetensors",
|
| 118 |
+
"model.language_model.layers.18.input_layernorm.weight": "model-00004-of-00014.safetensors",
|
| 119 |
+
"model.language_model.layers.18.mlp.down_proj.weight": "model-00005-of-00014.safetensors",
|
| 120 |
+
"model.language_model.layers.18.mlp.gate_proj.weight": "model-00004-of-00014.safetensors",
|
| 121 |
+
"model.language_model.layers.18.mlp.up_proj.weight": "model-00005-of-00014.safetensors",
|
| 122 |
+
"model.language_model.layers.18.post_attention_layernorm.weight": "model-00004-of-00014.safetensors",
|
| 123 |
+
"model.language_model.layers.18.self_attn.k_norm.weight": "model-00004-of-00014.safetensors",
|
| 124 |
+
"model.language_model.layers.18.self_attn.k_proj.weight": "model-00004-of-00014.safetensors",
|
| 125 |
+
"model.language_model.layers.18.self_attn.o_proj.weight": "model-00004-of-00014.safetensors",
|
| 126 |
+
"model.language_model.layers.18.self_attn.q_norm.weight": "model-00004-of-00014.safetensors",
|
| 127 |
+
"model.language_model.layers.18.self_attn.q_proj.weight": "model-00004-of-00014.safetensors",
|
| 128 |
+
"model.language_model.layers.18.self_attn.v_proj.weight": "model-00004-of-00014.safetensors",
|
| 129 |
+
"model.language_model.layers.19.input_layernorm.weight": "model-00005-of-00014.safetensors",
|
| 130 |
+
"model.language_model.layers.19.mlp.down_proj.weight": "model-00005-of-00014.safetensors",
|
| 131 |
+
"model.language_model.layers.19.mlp.gate_proj.weight": "model-00005-of-00014.safetensors",
|
| 132 |
+
"model.language_model.layers.19.mlp.up_proj.weight": "model-00005-of-00014.safetensors",
|
| 133 |
+
"model.language_model.layers.19.post_attention_layernorm.weight": "model-00005-of-00014.safetensors",
|
| 134 |
+
"model.language_model.layers.19.self_attn.k_norm.weight": "model-00005-of-00014.safetensors",
|
| 135 |
+
"model.language_model.layers.19.self_attn.k_proj.weight": "model-00005-of-00014.safetensors",
|
| 136 |
+
"model.language_model.layers.19.self_attn.o_proj.weight": "model-00005-of-00014.safetensors",
|
| 137 |
+
"model.language_model.layers.19.self_attn.q_norm.weight": "model-00005-of-00014.safetensors",
|
| 138 |
+
"model.language_model.layers.19.self_attn.q_proj.weight": "model-00005-of-00014.safetensors",
|
| 139 |
+
"model.language_model.layers.19.self_attn.v_proj.weight": "model-00005-of-00014.safetensors",
|
| 140 |
+
"model.language_model.layers.2.input_layernorm.weight": "model-00001-of-00014.safetensors",
|
| 141 |
+
"model.language_model.layers.2.mlp.down_proj.weight": "model-00001-of-00014.safetensors",
|
| 142 |
+
"model.language_model.layers.2.mlp.gate_proj.weight": "model-00001-of-00014.safetensors",
|
| 143 |
+
"model.language_model.layers.2.mlp.up_proj.weight": "model-00001-of-00014.safetensors",
|
| 144 |
+
"model.language_model.layers.2.post_attention_layernorm.weight": "model-00001-of-00014.safetensors",
|
| 145 |
+
"model.language_model.layers.2.self_attn.k_norm.weight": "model-00001-of-00014.safetensors",
|
| 146 |
+
"model.language_model.layers.2.self_attn.k_proj.weight": "model-00001-of-00014.safetensors",
|
| 147 |
+
"model.language_model.layers.2.self_attn.o_proj.weight": "model-00001-of-00014.safetensors",
|
| 148 |
+
"model.language_model.layers.2.self_attn.q_norm.weight": "model-00001-of-00014.safetensors",
|
| 149 |
+
"model.language_model.layers.2.self_attn.q_proj.weight": "model-00001-of-00014.safetensors",
|
| 150 |
+
"model.language_model.layers.2.self_attn.v_proj.weight": "model-00001-of-00014.safetensors",
|
| 151 |
+
"model.language_model.layers.20.input_layernorm.weight": "model-00005-of-00014.safetensors",
|
| 152 |
+
"model.language_model.layers.20.mlp.down_proj.weight": "model-00005-of-00014.safetensors",
|
| 153 |
+
"model.language_model.layers.20.mlp.gate_proj.weight": "model-00005-of-00014.safetensors",
|
| 154 |
+
"model.language_model.layers.20.mlp.up_proj.weight": "model-00005-of-00014.safetensors",
|
| 155 |
+
"model.language_model.layers.20.post_attention_layernorm.weight": "model-00005-of-00014.safetensors",
|
| 156 |
+
"model.language_model.layers.20.self_attn.k_norm.weight": "model-00005-of-00014.safetensors",
|
| 157 |
+
"model.language_model.layers.20.self_attn.k_proj.weight": "model-00005-of-00014.safetensors",
|
| 158 |
+
"model.language_model.layers.20.self_attn.o_proj.weight": "model-00005-of-00014.safetensors",
|
| 159 |
+
"model.language_model.layers.20.self_attn.q_norm.weight": "model-00005-of-00014.safetensors",
|
| 160 |
+
"model.language_model.layers.20.self_attn.q_proj.weight": "model-00005-of-00014.safetensors",
|
| 161 |
+
"model.language_model.layers.20.self_attn.v_proj.weight": "model-00005-of-00014.safetensors",
|
| 162 |
+
"model.language_model.layers.21.input_layernorm.weight": "model-00005-of-00014.safetensors",
|
| 163 |
+
"model.language_model.layers.21.mlp.down_proj.weight": "model-00005-of-00014.safetensors",
|
| 164 |
+
"model.language_model.layers.21.mlp.gate_proj.weight": "model-00005-of-00014.safetensors",
|
| 165 |
+
"model.language_model.layers.21.mlp.up_proj.weight": "model-00005-of-00014.safetensors",
|
| 166 |
+
"model.language_model.layers.21.post_attention_layernorm.weight": "model-00005-of-00014.safetensors",
|
| 167 |
+
"model.language_model.layers.21.self_attn.k_norm.weight": "model-00005-of-00014.safetensors",
|
| 168 |
+
"model.language_model.layers.21.self_attn.k_proj.weight": "model-00005-of-00014.safetensors",
|
| 169 |
+
"model.language_model.layers.21.self_attn.o_proj.weight": "model-00005-of-00014.safetensors",
|
| 170 |
+
"model.language_model.layers.21.self_attn.q_norm.weight": "model-00005-of-00014.safetensors",
|
| 171 |
+
"model.language_model.layers.21.self_attn.q_proj.weight": "model-00005-of-00014.safetensors",
|
| 172 |
+
"model.language_model.layers.21.self_attn.v_proj.weight": "model-00005-of-00014.safetensors",
|
| 173 |
+
"model.language_model.layers.22.input_layernorm.weight": "model-00005-of-00014.safetensors",
|
| 174 |
+
"model.language_model.layers.22.mlp.down_proj.weight": "model-00005-of-00014.safetensors",
|
| 175 |
+
"model.language_model.layers.22.mlp.gate_proj.weight": "model-00005-of-00014.safetensors",
|
| 176 |
+
"model.language_model.layers.22.mlp.up_proj.weight": "model-00005-of-00014.safetensors",
|
| 177 |
+
"model.language_model.layers.22.post_attention_layernorm.weight": "model-00005-of-00014.safetensors",
|
| 178 |
+
"model.language_model.layers.22.self_attn.k_norm.weight": "model-00005-of-00014.safetensors",
|
| 179 |
+
"model.language_model.layers.22.self_attn.k_proj.weight": "model-00005-of-00014.safetensors",
|
| 180 |
+
"model.language_model.layers.22.self_attn.o_proj.weight": "model-00005-of-00014.safetensors",
|
| 181 |
+
"model.language_model.layers.22.self_attn.q_norm.weight": "model-00005-of-00014.safetensors",
|
| 182 |
+
"model.language_model.layers.22.self_attn.q_proj.weight": "model-00005-of-00014.safetensors",
|
| 183 |
+
"model.language_model.layers.22.self_attn.v_proj.weight": "model-00005-of-00014.safetensors",
|
| 184 |
+
"model.language_model.layers.23.input_layernorm.weight": "model-00005-of-00014.safetensors",
|
| 185 |
+
"model.language_model.layers.23.mlp.down_proj.weight": "model-00006-of-00014.safetensors",
|
| 186 |
+
"model.language_model.layers.23.mlp.gate_proj.weight": "model-00005-of-00014.safetensors",
|
| 187 |
+
"model.language_model.layers.23.mlp.up_proj.weight": "model-00006-of-00014.safetensors",
|
| 188 |
+
"model.language_model.layers.23.post_attention_layernorm.weight": "model-00005-of-00014.safetensors",
|
| 189 |
+
"model.language_model.layers.23.self_attn.k_norm.weight": "model-00005-of-00014.safetensors",
|
| 190 |
+
"model.language_model.layers.23.self_attn.k_proj.weight": "model-00005-of-00014.safetensors",
|
| 191 |
+
"model.language_model.layers.23.self_attn.o_proj.weight": "model-00005-of-00014.safetensors",
|
| 192 |
+
"model.language_model.layers.23.self_attn.q_norm.weight": "model-00005-of-00014.safetensors",
|
| 193 |
+
"model.language_model.layers.23.self_attn.q_proj.weight": "model-00005-of-00014.safetensors",
|
| 194 |
+
"model.language_model.layers.23.self_attn.v_proj.weight": "model-00005-of-00014.safetensors",
|
| 195 |
+
"model.language_model.layers.24.input_layernorm.weight": "model-00006-of-00014.safetensors",
|
| 196 |
+
"model.language_model.layers.24.mlp.down_proj.weight": "model-00006-of-00014.safetensors",
|
| 197 |
+
"model.language_model.layers.24.mlp.gate_proj.weight": "model-00006-of-00014.safetensors",
|
| 198 |
+
"model.language_model.layers.24.mlp.up_proj.weight": "model-00006-of-00014.safetensors",
|
| 199 |
+
"model.language_model.layers.24.post_attention_layernorm.weight": "model-00006-of-00014.safetensors",
|
| 200 |
+
"model.language_model.layers.24.self_attn.k_norm.weight": "model-00006-of-00014.safetensors",
|
| 201 |
+
"model.language_model.layers.24.self_attn.k_proj.weight": "model-00006-of-00014.safetensors",
|
| 202 |
+
"model.language_model.layers.24.self_attn.o_proj.weight": "model-00006-of-00014.safetensors",
|
| 203 |
+
"model.language_model.layers.24.self_attn.q_norm.weight": "model-00006-of-00014.safetensors",
|
| 204 |
+
"model.language_model.layers.24.self_attn.q_proj.weight": "model-00006-of-00014.safetensors",
|
| 205 |
+
"model.language_model.layers.24.self_attn.v_proj.weight": "model-00006-of-00014.safetensors",
|
| 206 |
+
"model.language_model.layers.25.input_layernorm.weight": "model-00006-of-00014.safetensors",
|
| 207 |
+
"model.language_model.layers.25.mlp.down_proj.weight": "model-00006-of-00014.safetensors",
|
| 208 |
+
"model.language_model.layers.25.mlp.gate_proj.weight": "model-00006-of-00014.safetensors",
|
| 209 |
+
"model.language_model.layers.25.mlp.up_proj.weight": "model-00006-of-00014.safetensors",
|
| 210 |
+
"model.language_model.layers.25.post_attention_layernorm.weight": "model-00006-of-00014.safetensors",
|
| 211 |
+
"model.language_model.layers.25.self_attn.k_norm.weight": "model-00006-of-00014.safetensors",
|
| 212 |
+
"model.language_model.layers.25.self_attn.k_proj.weight": "model-00006-of-00014.safetensors",
|
| 213 |
+
"model.language_model.layers.25.self_attn.o_proj.weight": "model-00006-of-00014.safetensors",
|
| 214 |
+
"model.language_model.layers.25.self_attn.q_norm.weight": "model-00006-of-00014.safetensors",
|
| 215 |
+
"model.language_model.layers.25.self_attn.q_proj.weight": "model-00006-of-00014.safetensors",
|
| 216 |
+
"model.language_model.layers.25.self_attn.v_proj.weight": "model-00006-of-00014.safetensors",
|
| 217 |
+
"model.language_model.layers.26.input_layernorm.weight": "model-00006-of-00014.safetensors",
|
| 218 |
+
"model.language_model.layers.26.mlp.down_proj.weight": "model-00006-of-00014.safetensors",
|
| 219 |
+
"model.language_model.layers.26.mlp.gate_proj.weight": "model-00006-of-00014.safetensors",
|
| 220 |
+
"model.language_model.layers.26.mlp.up_proj.weight": "model-00006-of-00014.safetensors",
|
| 221 |
+
"model.language_model.layers.26.post_attention_layernorm.weight": "model-00006-of-00014.safetensors",
|
| 222 |
+
"model.language_model.layers.26.self_attn.k_norm.weight": "model-00006-of-00014.safetensors",
|
| 223 |
+
"model.language_model.layers.26.self_attn.k_proj.weight": "model-00006-of-00014.safetensors",
|
| 224 |
+
"model.language_model.layers.26.self_attn.o_proj.weight": "model-00006-of-00014.safetensors",
|
| 225 |
+
"model.language_model.layers.26.self_attn.q_norm.weight": "model-00006-of-00014.safetensors",
|
| 226 |
+
"model.language_model.layers.26.self_attn.q_proj.weight": "model-00006-of-00014.safetensors",
|
| 227 |
+
"model.language_model.layers.26.self_attn.v_proj.weight": "model-00006-of-00014.safetensors",
|
| 228 |
+
"model.language_model.layers.27.input_layernorm.weight": "model-00006-of-00014.safetensors",
|
| 229 |
+
"model.language_model.layers.27.mlp.down_proj.weight": "model-00006-of-00014.safetensors",
|
| 230 |
+
"model.language_model.layers.27.mlp.gate_proj.weight": "model-00006-of-00014.safetensors",
|
| 231 |
+
"model.language_model.layers.27.mlp.up_proj.weight": "model-00006-of-00014.safetensors",
|
| 232 |
+
"model.language_model.layers.27.post_attention_layernorm.weight": "model-00006-of-00014.safetensors",
|
| 233 |
+
"model.language_model.layers.27.self_attn.k_norm.weight": "model-00006-of-00014.safetensors",
|
| 234 |
+
"model.language_model.layers.27.self_attn.k_proj.weight": "model-00006-of-00014.safetensors",
|
| 235 |
+
"model.language_model.layers.27.self_attn.o_proj.weight": "model-00006-of-00014.safetensors",
|
| 236 |
+
"model.language_model.layers.27.self_attn.q_norm.weight": "model-00006-of-00014.safetensors",
|
| 237 |
+
"model.language_model.layers.27.self_attn.q_proj.weight": "model-00006-of-00014.safetensors",
|
| 238 |
+
"model.language_model.layers.27.self_attn.v_proj.weight": "model-00006-of-00014.safetensors",
|
| 239 |
+
"model.language_model.layers.28.input_layernorm.weight": "model-00006-of-00014.safetensors",
|
| 240 |
+
"model.language_model.layers.28.mlp.down_proj.weight": "model-00007-of-00014.safetensors",
|
| 241 |
+
"model.language_model.layers.28.mlp.gate_proj.weight": "model-00006-of-00014.safetensors",
|
| 242 |
+
"model.language_model.layers.28.mlp.up_proj.weight": "model-00007-of-00014.safetensors",
|
| 243 |
+
"model.language_model.layers.28.post_attention_layernorm.weight": "model-00006-of-00014.safetensors",
|
| 244 |
+
"model.language_model.layers.28.self_attn.k_norm.weight": "model-00006-of-00014.safetensors",
|
| 245 |
+
"model.language_model.layers.28.self_attn.k_proj.weight": "model-00006-of-00014.safetensors",
|
| 246 |
+
"model.language_model.layers.28.self_attn.o_proj.weight": "model-00006-of-00014.safetensors",
|
| 247 |
+
"model.language_model.layers.28.self_attn.q_norm.weight": "model-00006-of-00014.safetensors",
|
| 248 |
+
"model.language_model.layers.28.self_attn.q_proj.weight": "model-00006-of-00014.safetensors",
|
| 249 |
+
"model.language_model.layers.28.self_attn.v_proj.weight": "model-00006-of-00014.safetensors",
|
| 250 |
+
"model.language_model.layers.29.input_layernorm.weight": "model-00007-of-00014.safetensors",
|
| 251 |
+
"model.language_model.layers.29.mlp.down_proj.weight": "model-00007-of-00014.safetensors",
|
| 252 |
+
"model.language_model.layers.29.mlp.gate_proj.weight": "model-00007-of-00014.safetensors",
|
| 253 |
+
"model.language_model.layers.29.mlp.up_proj.weight": "model-00007-of-00014.safetensors",
|
| 254 |
+
"model.language_model.layers.29.post_attention_layernorm.weight": "model-00007-of-00014.safetensors",
|
| 255 |
+
"model.language_model.layers.29.self_attn.k_norm.weight": "model-00007-of-00014.safetensors",
|
| 256 |
+
"model.language_model.layers.29.self_attn.k_proj.weight": "model-00007-of-00014.safetensors",
|
| 257 |
+
"model.language_model.layers.29.self_attn.o_proj.weight": "model-00007-of-00014.safetensors",
|
| 258 |
+
"model.language_model.layers.29.self_attn.q_norm.weight": "model-00007-of-00014.safetensors",
|
| 259 |
+
"model.language_model.layers.29.self_attn.q_proj.weight": "model-00007-of-00014.safetensors",
|
| 260 |
+
"model.language_model.layers.29.self_attn.v_proj.weight": "model-00007-of-00014.safetensors",
|
| 261 |
+
"model.language_model.layers.3.input_layernorm.weight": "model-00001-of-00014.safetensors",
|
| 262 |
+
"model.language_model.layers.3.mlp.down_proj.weight": "model-00002-of-00014.safetensors",
|
| 263 |
+
"model.language_model.layers.3.mlp.gate_proj.weight": "model-00001-of-00014.safetensors",
|
| 264 |
+
"model.language_model.layers.3.mlp.up_proj.weight": "model-00002-of-00014.safetensors",
|
| 265 |
+
"model.language_model.layers.3.post_attention_layernorm.weight": "model-00001-of-00014.safetensors",
|
| 266 |
+
"model.language_model.layers.3.self_attn.k_norm.weight": "model-00001-of-00014.safetensors",
|
| 267 |
+
"model.language_model.layers.3.self_attn.k_proj.weight": "model-00001-of-00014.safetensors",
|
| 268 |
+
"model.language_model.layers.3.self_attn.o_proj.weight": "model-00001-of-00014.safetensors",
|
| 269 |
+
"model.language_model.layers.3.self_attn.q_norm.weight": "model-00001-of-00014.safetensors",
|
| 270 |
+
"model.language_model.layers.3.self_attn.q_proj.weight": "model-00001-of-00014.safetensors",
|
| 271 |
+
"model.language_model.layers.3.self_attn.v_proj.weight": "model-00001-of-00014.safetensors",
|
| 272 |
+
"model.language_model.layers.30.input_layernorm.weight": "model-00007-of-00014.safetensors",
|
| 273 |
+
"model.language_model.layers.30.mlp.down_proj.weight": "model-00007-of-00014.safetensors",
|
| 274 |
+
"model.language_model.layers.30.mlp.gate_proj.weight": "model-00007-of-00014.safetensors",
|
| 275 |
+
"model.language_model.layers.30.mlp.up_proj.weight": "model-00007-of-00014.safetensors",
|
| 276 |
+
"model.language_model.layers.30.post_attention_layernorm.weight": "model-00007-of-00014.safetensors",
|
| 277 |
+
"model.language_model.layers.30.self_attn.k_norm.weight": "model-00007-of-00014.safetensors",
|
| 278 |
+
"model.language_model.layers.30.self_attn.k_proj.weight": "model-00007-of-00014.safetensors",
|
| 279 |
+
"model.language_model.layers.30.self_attn.o_proj.weight": "model-00007-of-00014.safetensors",
|
| 280 |
+
"model.language_model.layers.30.self_attn.q_norm.weight": "model-00007-of-00014.safetensors",
|
| 281 |
+
"model.language_model.layers.30.self_attn.q_proj.weight": "model-00007-of-00014.safetensors",
|
| 282 |
+
"model.language_model.layers.30.self_attn.v_proj.weight": "model-00007-of-00014.safetensors",
|
| 283 |
+
"model.language_model.layers.31.input_layernorm.weight": "model-00007-of-00014.safetensors",
|
| 284 |
+
"model.language_model.layers.31.mlp.down_proj.weight": "model-00007-of-00014.safetensors",
|
| 285 |
+
"model.language_model.layers.31.mlp.gate_proj.weight": "model-00007-of-00014.safetensors",
|
| 286 |
+
"model.language_model.layers.31.mlp.up_proj.weight": "model-00007-of-00014.safetensors",
|
| 287 |
+
"model.language_model.layers.31.post_attention_layernorm.weight": "model-00007-of-00014.safetensors",
|
| 288 |
+
"model.language_model.layers.31.self_attn.k_norm.weight": "model-00007-of-00014.safetensors",
|
| 289 |
+
"model.language_model.layers.31.self_attn.k_proj.weight": "model-00007-of-00014.safetensors",
|
| 290 |
+
"model.language_model.layers.31.self_attn.o_proj.weight": "model-00007-of-00014.safetensors",
|
| 291 |
+
"model.language_model.layers.31.self_attn.q_norm.weight": "model-00007-of-00014.safetensors",
|
| 292 |
+
"model.language_model.layers.31.self_attn.q_proj.weight": "model-00007-of-00014.safetensors",
|
| 293 |
+
"model.language_model.layers.31.self_attn.v_proj.weight": "model-00007-of-00014.safetensors",
|
| 294 |
+
"model.language_model.layers.32.input_layernorm.weight": "model-00007-of-00014.safetensors",
|
| 295 |
+
"model.language_model.layers.32.mlp.down_proj.weight": "model-00007-of-00014.safetensors",
|
| 296 |
+
"model.language_model.layers.32.mlp.gate_proj.weight": "model-00007-of-00014.safetensors",
|
| 297 |
+
"model.language_model.layers.32.mlp.up_proj.weight": "model-00007-of-00014.safetensors",
|
| 298 |
+
"model.language_model.layers.32.post_attention_layernorm.weight": "model-00007-of-00014.safetensors",
|
| 299 |
+
"model.language_model.layers.32.self_attn.k_norm.weight": "model-00007-of-00014.safetensors",
|
| 300 |
+
"model.language_model.layers.32.self_attn.k_proj.weight": "model-00007-of-00014.safetensors",
|
| 301 |
+
"model.language_model.layers.32.self_attn.o_proj.weight": "model-00007-of-00014.safetensors",
|
| 302 |
+
"model.language_model.layers.32.self_attn.q_norm.weight": "model-00007-of-00014.safetensors",
|
| 303 |
+
"model.language_model.layers.32.self_attn.q_proj.weight": "model-00007-of-00014.safetensors",
|
| 304 |
+
"model.language_model.layers.32.self_attn.v_proj.weight": "model-00007-of-00014.safetensors",
|
| 305 |
+
"model.language_model.layers.33.input_layernorm.weight": "model-00007-of-00014.safetensors",
|
| 306 |
+
"model.language_model.layers.33.mlp.down_proj.weight": "model-00008-of-00014.safetensors",
|
| 307 |
+
"model.language_model.layers.33.mlp.gate_proj.weight": "model-00007-of-00014.safetensors",
|
| 308 |
+
"model.language_model.layers.33.mlp.up_proj.weight": "model-00008-of-00014.safetensors",
|
| 309 |
+
"model.language_model.layers.33.post_attention_layernorm.weight": "model-00007-of-00014.safetensors",
|
| 310 |
+
"model.language_model.layers.33.self_attn.k_norm.weight": "model-00007-of-00014.safetensors",
|
| 311 |
+
"model.language_model.layers.33.self_attn.k_proj.weight": "model-00007-of-00014.safetensors",
|
| 312 |
+
"model.language_model.layers.33.self_attn.o_proj.weight": "model-00007-of-00014.safetensors",
|
| 313 |
+
"model.language_model.layers.33.self_attn.q_norm.weight": "model-00007-of-00014.safetensors",
|
| 314 |
+
"model.language_model.layers.33.self_attn.q_proj.weight": "model-00007-of-00014.safetensors",
|
| 315 |
+
"model.language_model.layers.33.self_attn.v_proj.weight": "model-00007-of-00014.safetensors",
|
| 316 |
+
"model.language_model.layers.34.input_layernorm.weight": "model-00008-of-00014.safetensors",
|
| 317 |
+
"model.language_model.layers.34.mlp.down_proj.weight": "model-00008-of-00014.safetensors",
|
| 318 |
+
"model.language_model.layers.34.mlp.gate_proj.weight": "model-00008-of-00014.safetensors",
|
| 319 |
+
"model.language_model.layers.34.mlp.up_proj.weight": "model-00008-of-00014.safetensors",
|
| 320 |
+
"model.language_model.layers.34.post_attention_layernorm.weight": "model-00008-of-00014.safetensors",
|
| 321 |
+
"model.language_model.layers.34.self_attn.k_norm.weight": "model-00008-of-00014.safetensors",
|
| 322 |
+
"model.language_model.layers.34.self_attn.k_proj.weight": "model-00008-of-00014.safetensors",
|
| 323 |
+
"model.language_model.layers.34.self_attn.o_proj.weight": "model-00008-of-00014.safetensors",
|
| 324 |
+
"model.language_model.layers.34.self_attn.q_norm.weight": "model-00008-of-00014.safetensors",
|
| 325 |
+
"model.language_model.layers.34.self_attn.q_proj.weight": "model-00008-of-00014.safetensors",
|
| 326 |
+
"model.language_model.layers.34.self_attn.v_proj.weight": "model-00008-of-00014.safetensors",
|
| 327 |
+
"model.language_model.layers.35.input_layernorm.weight": "model-00008-of-00014.safetensors",
|
| 328 |
+
"model.language_model.layers.35.mlp.down_proj.weight": "model-00008-of-00014.safetensors",
|
| 329 |
+
"model.language_model.layers.35.mlp.gate_proj.weight": "model-00008-of-00014.safetensors",
|
| 330 |
+
"model.language_model.layers.35.mlp.up_proj.weight": "model-00008-of-00014.safetensors",
|
| 331 |
+
"model.language_model.layers.35.post_attention_layernorm.weight": "model-00008-of-00014.safetensors",
|
| 332 |
+
"model.language_model.layers.35.self_attn.k_norm.weight": "model-00008-of-00014.safetensors",
|
| 333 |
+
"model.language_model.layers.35.self_attn.k_proj.weight": "model-00008-of-00014.safetensors",
|
| 334 |
+
"model.language_model.layers.35.self_attn.o_proj.weight": "model-00008-of-00014.safetensors",
|
| 335 |
+
"model.language_model.layers.35.self_attn.q_norm.weight": "model-00008-of-00014.safetensors",
|
| 336 |
+
"model.language_model.layers.35.self_attn.q_proj.weight": "model-00008-of-00014.safetensors",
|
| 337 |
+
"model.language_model.layers.35.self_attn.v_proj.weight": "model-00008-of-00014.safetensors",
|
| 338 |
+
"model.language_model.layers.36.input_layernorm.weight": "model-00008-of-00014.safetensors",
|
| 339 |
+
"model.language_model.layers.36.mlp.down_proj.weight": "model-00008-of-00014.safetensors",
|
| 340 |
+
"model.language_model.layers.36.mlp.gate_proj.weight": "model-00008-of-00014.safetensors",
|
| 341 |
+
"model.language_model.layers.36.mlp.up_proj.weight": "model-00008-of-00014.safetensors",
|
| 342 |
+
"model.language_model.layers.36.post_attention_layernorm.weight": "model-00008-of-00014.safetensors",
|
| 343 |
+
"model.language_model.layers.36.self_attn.k_norm.weight": "model-00008-of-00014.safetensors",
|
| 344 |
+
"model.language_model.layers.36.self_attn.k_proj.weight": "model-00008-of-00014.safetensors",
|
| 345 |
+
"model.language_model.layers.36.self_attn.o_proj.weight": "model-00008-of-00014.safetensors",
|
| 346 |
+
"model.language_model.layers.36.self_attn.q_norm.weight": "model-00008-of-00014.safetensors",
|
| 347 |
+
"model.language_model.layers.36.self_attn.q_proj.weight": "model-00008-of-00014.safetensors",
|
| 348 |
+
"model.language_model.layers.36.self_attn.v_proj.weight": "model-00008-of-00014.safetensors",
|
| 349 |
+
"model.language_model.layers.37.input_layernorm.weight": "model-00008-of-00014.safetensors",
|
| 350 |
+
"model.language_model.layers.37.mlp.down_proj.weight": "model-00008-of-00014.safetensors",
|
| 351 |
+
"model.language_model.layers.37.mlp.gate_proj.weight": "model-00008-of-00014.safetensors",
|
| 352 |
+
"model.language_model.layers.37.mlp.up_proj.weight": "model-00008-of-00014.safetensors",
|
| 353 |
+
"model.language_model.layers.37.post_attention_layernorm.weight": "model-00008-of-00014.safetensors",
|
| 354 |
+
"model.language_model.layers.37.self_attn.k_norm.weight": "model-00008-of-00014.safetensors",
|
| 355 |
+
"model.language_model.layers.37.self_attn.k_proj.weight": "model-00008-of-00014.safetensors",
|
| 356 |
+
"model.language_model.layers.37.self_attn.o_proj.weight": "model-00008-of-00014.safetensors",
|
| 357 |
+
"model.language_model.layers.37.self_attn.q_norm.weight": "model-00008-of-00014.safetensors",
|
| 358 |
+
"model.language_model.layers.37.self_attn.q_proj.weight": "model-00008-of-00014.safetensors",
|
| 359 |
+
"model.language_model.layers.37.self_attn.v_proj.weight": "model-00008-of-00014.safetensors",
|
| 360 |
+
"model.language_model.layers.38.input_layernorm.weight": "model-00008-of-00014.safetensors",
|
| 361 |
+
"model.language_model.layers.38.mlp.down_proj.weight": "model-00009-of-00014.safetensors",
|
| 362 |
+
"model.language_model.layers.38.mlp.gate_proj.weight": "model-00008-of-00014.safetensors",
|
| 363 |
+
"model.language_model.layers.38.mlp.up_proj.weight": "model-00009-of-00014.safetensors",
|
| 364 |
+
"model.language_model.layers.38.post_attention_layernorm.weight": "model-00008-of-00014.safetensors",
|
| 365 |
+
"model.language_model.layers.38.self_attn.k_norm.weight": "model-00008-of-00014.safetensors",
|
| 366 |
+
"model.language_model.layers.38.self_attn.k_proj.weight": "model-00008-of-00014.safetensors",
|
| 367 |
+
"model.language_model.layers.38.self_attn.o_proj.weight": "model-00008-of-00014.safetensors",
|
| 368 |
+
"model.language_model.layers.38.self_attn.q_norm.weight": "model-00008-of-00014.safetensors",
|
| 369 |
+
"model.language_model.layers.38.self_attn.q_proj.weight": "model-00008-of-00014.safetensors",
|
| 370 |
+
"model.language_model.layers.38.self_attn.v_proj.weight": "model-00008-of-00014.safetensors",
|
| 371 |
+
"model.language_model.layers.39.input_layernorm.weight": "model-00009-of-00014.safetensors",
|
| 372 |
+
"model.language_model.layers.39.mlp.down_proj.weight": "model-00009-of-00014.safetensors",
|
| 373 |
+
"model.language_model.layers.39.mlp.gate_proj.weight": "model-00009-of-00014.safetensors",
|
| 374 |
+
"model.language_model.layers.39.mlp.up_proj.weight": "model-00009-of-00014.safetensors",
|
| 375 |
+
"model.language_model.layers.39.post_attention_layernorm.weight": "model-00009-of-00014.safetensors",
|
| 376 |
+
"model.language_model.layers.39.self_attn.k_norm.weight": "model-00009-of-00014.safetensors",
|
| 377 |
+
"model.language_model.layers.39.self_attn.k_proj.weight": "model-00009-of-00014.safetensors",
|
| 378 |
+
"model.language_model.layers.39.self_attn.o_proj.weight": "model-00009-of-00014.safetensors",
|
| 379 |
+
"model.language_model.layers.39.self_attn.q_norm.weight": "model-00009-of-00014.safetensors",
|
| 380 |
+
"model.language_model.layers.39.self_attn.q_proj.weight": "model-00009-of-00014.safetensors",
|
| 381 |
+
"model.language_model.layers.39.self_attn.v_proj.weight": "model-00009-of-00014.safetensors",
|
| 382 |
+
"model.language_model.layers.4.input_layernorm.weight": "model-00002-of-00014.safetensors",
|
| 383 |
+
"model.language_model.layers.4.mlp.down_proj.weight": "model-00002-of-00014.safetensors",
|
| 384 |
+
"model.language_model.layers.4.mlp.gate_proj.weight": "model-00002-of-00014.safetensors",
|
| 385 |
+
"model.language_model.layers.4.mlp.up_proj.weight": "model-00002-of-00014.safetensors",
|
| 386 |
+
"model.language_model.layers.4.post_attention_layernorm.weight": "model-00002-of-00014.safetensors",
|
| 387 |
+
"model.language_model.layers.4.self_attn.k_norm.weight": "model-00002-of-00014.safetensors",
|
| 388 |
+
"model.language_model.layers.4.self_attn.k_proj.weight": "model-00002-of-00014.safetensors",
|
| 389 |
+
"model.language_model.layers.4.self_attn.o_proj.weight": "model-00002-of-00014.safetensors",
|
| 390 |
+
"model.language_model.layers.4.self_attn.q_norm.weight": "model-00002-of-00014.safetensors",
|
| 391 |
+
"model.language_model.layers.4.self_attn.q_proj.weight": "model-00002-of-00014.safetensors",
|
| 392 |
+
"model.language_model.layers.4.self_attn.v_proj.weight": "model-00002-of-00014.safetensors",
|
| 393 |
+
"model.language_model.layers.40.input_layernorm.weight": "model-00009-of-00014.safetensors",
|
| 394 |
+
"model.language_model.layers.40.mlp.down_proj.weight": "model-00009-of-00014.safetensors",
|
| 395 |
+
"model.language_model.layers.40.mlp.gate_proj.weight": "model-00009-of-00014.safetensors",
|
| 396 |
+
"model.language_model.layers.40.mlp.up_proj.weight": "model-00009-of-00014.safetensors",
|
| 397 |
+
"model.language_model.layers.40.post_attention_layernorm.weight": "model-00009-of-00014.safetensors",
|
| 398 |
+
"model.language_model.layers.40.self_attn.k_norm.weight": "model-00009-of-00014.safetensors",
|
| 399 |
+
"model.language_model.layers.40.self_attn.k_proj.weight": "model-00009-of-00014.safetensors",
|
| 400 |
+
"model.language_model.layers.40.self_attn.o_proj.weight": "model-00009-of-00014.safetensors",
|
| 401 |
+
"model.language_model.layers.40.self_attn.q_norm.weight": "model-00009-of-00014.safetensors",
|
| 402 |
+
"model.language_model.layers.40.self_attn.q_proj.weight": "model-00009-of-00014.safetensors",
|
| 403 |
+
"model.language_model.layers.40.self_attn.v_proj.weight": "model-00009-of-00014.safetensors",
|
| 404 |
+
"model.language_model.layers.41.input_layernorm.weight": "model-00009-of-00014.safetensors",
|
| 405 |
+
"model.language_model.layers.41.mlp.down_proj.weight": "model-00009-of-00014.safetensors",
|
| 406 |
+
"model.language_model.layers.41.mlp.gate_proj.weight": "model-00009-of-00014.safetensors",
|
| 407 |
+
"model.language_model.layers.41.mlp.up_proj.weight": "model-00009-of-00014.safetensors",
|
| 408 |
+
"model.language_model.layers.41.post_attention_layernorm.weight": "model-00009-of-00014.safetensors",
|
| 409 |
+
"model.language_model.layers.41.self_attn.k_norm.weight": "model-00009-of-00014.safetensors",
|
| 410 |
+
"model.language_model.layers.41.self_attn.k_proj.weight": "model-00009-of-00014.safetensors",
|
| 411 |
+
"model.language_model.layers.41.self_attn.o_proj.weight": "model-00009-of-00014.safetensors",
|
| 412 |
+
"model.language_model.layers.41.self_attn.q_norm.weight": "model-00009-of-00014.safetensors",
|
| 413 |
+
"model.language_model.layers.41.self_attn.q_proj.weight": "model-00009-of-00014.safetensors",
|
| 414 |
+
"model.language_model.layers.41.self_attn.v_proj.weight": "model-00009-of-00014.safetensors",
|
| 415 |
+
"model.language_model.layers.42.input_layernorm.weight": "model-00009-of-00014.safetensors",
|
| 416 |
+
"model.language_model.layers.42.mlp.down_proj.weight": "model-00009-of-00014.safetensors",
|
| 417 |
+
"model.language_model.layers.42.mlp.gate_proj.weight": "model-00009-of-00014.safetensors",
|
| 418 |
+
"model.language_model.layers.42.mlp.up_proj.weight": "model-00009-of-00014.safetensors",
|
| 419 |
+
"model.language_model.layers.42.post_attention_layernorm.weight": "model-00009-of-00014.safetensors",
|
| 420 |
+
"model.language_model.layers.42.self_attn.k_norm.weight": "model-00009-of-00014.safetensors",
|
| 421 |
+
"model.language_model.layers.42.self_attn.k_proj.weight": "model-00009-of-00014.safetensors",
|
| 422 |
+
"model.language_model.layers.42.self_attn.o_proj.weight": "model-00009-of-00014.safetensors",
|
| 423 |
+
"model.language_model.layers.42.self_attn.q_norm.weight": "model-00009-of-00014.safetensors",
|
| 424 |
+
"model.language_model.layers.42.self_attn.q_proj.weight": "model-00009-of-00014.safetensors",
|
| 425 |
+
"model.language_model.layers.42.self_attn.v_proj.weight": "model-00009-of-00014.safetensors",
|
| 426 |
+
"model.language_model.layers.43.input_layernorm.weight": "model-00009-of-00014.safetensors",
|
| 427 |
+
"model.language_model.layers.43.mlp.down_proj.weight": "model-00010-of-00014.safetensors",
|
| 428 |
+
"model.language_model.layers.43.mlp.gate_proj.weight": "model-00009-of-00014.safetensors",
|
| 429 |
+
"model.language_model.layers.43.mlp.up_proj.weight": "model-00010-of-00014.safetensors",
|
| 430 |
+
"model.language_model.layers.43.post_attention_layernorm.weight": "model-00009-of-00014.safetensors",
|
| 431 |
+
"model.language_model.layers.43.self_attn.k_norm.weight": "model-00009-of-00014.safetensors",
|
| 432 |
+
"model.language_model.layers.43.self_attn.k_proj.weight": "model-00009-of-00014.safetensors",
|
| 433 |
+
"model.language_model.layers.43.self_attn.o_proj.weight": "model-00009-of-00014.safetensors",
|
| 434 |
+
"model.language_model.layers.43.self_attn.q_norm.weight": "model-00009-of-00014.safetensors",
|
| 435 |
+
"model.language_model.layers.43.self_attn.q_proj.weight": "model-00009-of-00014.safetensors",
|
| 436 |
+
"model.language_model.layers.43.self_attn.v_proj.weight": "model-00009-of-00014.safetensors",
|
| 437 |
+
"model.language_model.layers.44.input_layernorm.weight": "model-00010-of-00014.safetensors",
|
| 438 |
+
"model.language_model.layers.44.mlp.down_proj.weight": "model-00010-of-00014.safetensors",
|
| 439 |
+
"model.language_model.layers.44.mlp.gate_proj.weight": "model-00010-of-00014.safetensors",
|
| 440 |
+
"model.language_model.layers.44.mlp.up_proj.weight": "model-00010-of-00014.safetensors",
|
| 441 |
+
"model.language_model.layers.44.post_attention_layernorm.weight": "model-00010-of-00014.safetensors",
|
| 442 |
+
"model.language_model.layers.44.self_attn.k_norm.weight": "model-00010-of-00014.safetensors",
|
| 443 |
+
"model.language_model.layers.44.self_attn.k_proj.weight": "model-00010-of-00014.safetensors",
|
| 444 |
+
"model.language_model.layers.44.self_attn.o_proj.weight": "model-00010-of-00014.safetensors",
|
| 445 |
+
"model.language_model.layers.44.self_attn.q_norm.weight": "model-00010-of-00014.safetensors",
|
| 446 |
+
"model.language_model.layers.44.self_attn.q_proj.weight": "model-00010-of-00014.safetensors",
|
| 447 |
+
"model.language_model.layers.44.self_attn.v_proj.weight": "model-00010-of-00014.safetensors",
|
| 448 |
+
"model.language_model.layers.45.input_layernorm.weight": "model-00010-of-00014.safetensors",
|
| 449 |
+
"model.language_model.layers.45.mlp.down_proj.weight": "model-00010-of-00014.safetensors",
|
| 450 |
+
"model.language_model.layers.45.mlp.gate_proj.weight": "model-00010-of-00014.safetensors",
|
| 451 |
+
"model.language_model.layers.45.mlp.up_proj.weight": "model-00010-of-00014.safetensors",
|
| 452 |
+
"model.language_model.layers.45.post_attention_layernorm.weight": "model-00010-of-00014.safetensors",
|
| 453 |
+
"model.language_model.layers.45.self_attn.k_norm.weight": "model-00010-of-00014.safetensors",
|
| 454 |
+
"model.language_model.layers.45.self_attn.k_proj.weight": "model-00010-of-00014.safetensors",
|
| 455 |
+
"model.language_model.layers.45.self_attn.o_proj.weight": "model-00010-of-00014.safetensors",
|
| 456 |
+
"model.language_model.layers.45.self_attn.q_norm.weight": "model-00010-of-00014.safetensors",
|
| 457 |
+
"model.language_model.layers.45.self_attn.q_proj.weight": "model-00010-of-00014.safetensors",
|
| 458 |
+
"model.language_model.layers.45.self_attn.v_proj.weight": "model-00010-of-00014.safetensors",
|
| 459 |
+
"model.language_model.layers.46.input_layernorm.weight": "model-00010-of-00014.safetensors",
|
| 460 |
+
"model.language_model.layers.46.mlp.down_proj.weight": "model-00010-of-00014.safetensors",
|
| 461 |
+
"model.language_model.layers.46.mlp.gate_proj.weight": "model-00010-of-00014.safetensors",
|
| 462 |
+
"model.language_model.layers.46.mlp.up_proj.weight": "model-00010-of-00014.safetensors",
|
| 463 |
+
"model.language_model.layers.46.post_attention_layernorm.weight": "model-00010-of-00014.safetensors",
|
| 464 |
+
"model.language_model.layers.46.self_attn.k_norm.weight": "model-00010-of-00014.safetensors",
|
| 465 |
+
"model.language_model.layers.46.self_attn.k_proj.weight": "model-00010-of-00014.safetensors",
|
| 466 |
+
"model.language_model.layers.46.self_attn.o_proj.weight": "model-00010-of-00014.safetensors",
|
| 467 |
+
"model.language_model.layers.46.self_attn.q_norm.weight": "model-00010-of-00014.safetensors",
|
| 468 |
+
"model.language_model.layers.46.self_attn.q_proj.weight": "model-00010-of-00014.safetensors",
|
| 469 |
+
"model.language_model.layers.46.self_attn.v_proj.weight": "model-00010-of-00014.safetensors",
|
| 470 |
+
"model.language_model.layers.47.input_layernorm.weight": "model-00010-of-00014.safetensors",
|
| 471 |
+
"model.language_model.layers.47.mlp.down_proj.weight": "model-00010-of-00014.safetensors",
|
| 472 |
+
"model.language_model.layers.47.mlp.gate_proj.weight": "model-00010-of-00014.safetensors",
|
| 473 |
+
"model.language_model.layers.47.mlp.up_proj.weight": "model-00010-of-00014.safetensors",
|
| 474 |
+
"model.language_model.layers.47.post_attention_layernorm.weight": "model-00010-of-00014.safetensors",
|
| 475 |
+
"model.language_model.layers.47.self_attn.k_norm.weight": "model-00010-of-00014.safetensors",
|
| 476 |
+
"model.language_model.layers.47.self_attn.k_proj.weight": "model-00010-of-00014.safetensors",
|
| 477 |
+
"model.language_model.layers.47.self_attn.o_proj.weight": "model-00010-of-00014.safetensors",
|
| 478 |
+
"model.language_model.layers.47.self_attn.q_norm.weight": "model-00010-of-00014.safetensors",
|
| 479 |
+
"model.language_model.layers.47.self_attn.q_proj.weight": "model-00010-of-00014.safetensors",
|
| 480 |
+
"model.language_model.layers.47.self_attn.v_proj.weight": "model-00010-of-00014.safetensors",
|
| 481 |
+
"model.language_model.layers.48.input_layernorm.weight": "model-00010-of-00014.safetensors",
|
| 482 |
+
"model.language_model.layers.48.mlp.down_proj.weight": "model-00011-of-00014.safetensors",
|
| 483 |
+
"model.language_model.layers.48.mlp.gate_proj.weight": "model-00010-of-00014.safetensors",
|
| 484 |
+
"model.language_model.layers.48.mlp.up_proj.weight": "model-00011-of-00014.safetensors",
|
| 485 |
+
"model.language_model.layers.48.post_attention_layernorm.weight": "model-00010-of-00014.safetensors",
|
| 486 |
+
"model.language_model.layers.48.self_attn.k_norm.weight": "model-00010-of-00014.safetensors",
|
| 487 |
+
"model.language_model.layers.48.self_attn.k_proj.weight": "model-00010-of-00014.safetensors",
|
| 488 |
+
"model.language_model.layers.48.self_attn.o_proj.weight": "model-00010-of-00014.safetensors",
|
| 489 |
+
"model.language_model.layers.48.self_attn.q_norm.weight": "model-00010-of-00014.safetensors",
|
| 490 |
+
"model.language_model.layers.48.self_attn.q_proj.weight": "model-00010-of-00014.safetensors",
|
| 491 |
+
"model.language_model.layers.48.self_attn.v_proj.weight": "model-00010-of-00014.safetensors",
|
| 492 |
+
"model.language_model.layers.49.input_layernorm.weight": "model-00011-of-00014.safetensors",
|
| 493 |
+
"model.language_model.layers.49.mlp.down_proj.weight": "model-00011-of-00014.safetensors",
|
| 494 |
+
"model.language_model.layers.49.mlp.gate_proj.weight": "model-00011-of-00014.safetensors",
|
| 495 |
+
"model.language_model.layers.49.mlp.up_proj.weight": "model-00011-of-00014.safetensors",
|
| 496 |
+
"model.language_model.layers.49.post_attention_layernorm.weight": "model-00011-of-00014.safetensors",
|
| 497 |
+
"model.language_model.layers.49.self_attn.k_norm.weight": "model-00011-of-00014.safetensors",
|
| 498 |
+
"model.language_model.layers.49.self_attn.k_proj.weight": "model-00011-of-00014.safetensors",
|
| 499 |
+
"model.language_model.layers.49.self_attn.o_proj.weight": "model-00011-of-00014.safetensors",
|
| 500 |
+
"model.language_model.layers.49.self_attn.q_norm.weight": "model-00011-of-00014.safetensors",
|
| 501 |
+
"model.language_model.layers.49.self_attn.q_proj.weight": "model-00011-of-00014.safetensors",
|
| 502 |
+
"model.language_model.layers.49.self_attn.v_proj.weight": "model-00011-of-00014.safetensors",
|
| 503 |
+
"model.language_model.layers.5.input_layernorm.weight": "model-00002-of-00014.safetensors",
|
| 504 |
+
"model.language_model.layers.5.mlp.down_proj.weight": "model-00002-of-00014.safetensors",
|
| 505 |
+
"model.language_model.layers.5.mlp.gate_proj.weight": "model-00002-of-00014.safetensors",
|
| 506 |
+
"model.language_model.layers.5.mlp.up_proj.weight": "model-00002-of-00014.safetensors",
|
| 507 |
+
"model.language_model.layers.5.post_attention_layernorm.weight": "model-00002-of-00014.safetensors",
|
| 508 |
+
"model.language_model.layers.5.self_attn.k_norm.weight": "model-00002-of-00014.safetensors",
|
| 509 |
+
"model.language_model.layers.5.self_attn.k_proj.weight": "model-00002-of-00014.safetensors",
|
| 510 |
+
"model.language_model.layers.5.self_attn.o_proj.weight": "model-00002-of-00014.safetensors",
|
| 511 |
+
"model.language_model.layers.5.self_attn.q_norm.weight": "model-00002-of-00014.safetensors",
|
| 512 |
+
"model.language_model.layers.5.self_attn.q_proj.weight": "model-00002-of-00014.safetensors",
|
| 513 |
+
"model.language_model.layers.5.self_attn.v_proj.weight": "model-00002-of-00014.safetensors",
|
| 514 |
+
"model.language_model.layers.50.input_layernorm.weight": "model-00011-of-00014.safetensors",
|
| 515 |
+
"model.language_model.layers.50.mlp.down_proj.weight": "model-00011-of-00014.safetensors",
|
| 516 |
+
"model.language_model.layers.50.mlp.gate_proj.weight": "model-00011-of-00014.safetensors",
|
| 517 |
+
"model.language_model.layers.50.mlp.up_proj.weight": "model-00011-of-00014.safetensors",
|
| 518 |
+
"model.language_model.layers.50.post_attention_layernorm.weight": "model-00011-of-00014.safetensors",
|
| 519 |
+
"model.language_model.layers.50.self_attn.k_norm.weight": "model-00011-of-00014.safetensors",
|
| 520 |
+
"model.language_model.layers.50.self_attn.k_proj.weight": "model-00011-of-00014.safetensors",
|
| 521 |
+
"model.language_model.layers.50.self_attn.o_proj.weight": "model-00011-of-00014.safetensors",
|
| 522 |
+
"model.language_model.layers.50.self_attn.q_norm.weight": "model-00011-of-00014.safetensors",
|
| 523 |
+
"model.language_model.layers.50.self_attn.q_proj.weight": "model-00011-of-00014.safetensors",
|
| 524 |
+
"model.language_model.layers.50.self_attn.v_proj.weight": "model-00011-of-00014.safetensors",
|
| 525 |
+
"model.language_model.layers.51.input_layernorm.weight": "model-00011-of-00014.safetensors",
|
| 526 |
+
"model.language_model.layers.51.mlp.down_proj.weight": "model-00011-of-00014.safetensors",
|
| 527 |
+
"model.language_model.layers.51.mlp.gate_proj.weight": "model-00011-of-00014.safetensors",
|
| 528 |
+
"model.language_model.layers.51.mlp.up_proj.weight": "model-00011-of-00014.safetensors",
|
| 529 |
+
"model.language_model.layers.51.post_attention_layernorm.weight": "model-00011-of-00014.safetensors",
|
| 530 |
+
"model.language_model.layers.51.self_attn.k_norm.weight": "model-00011-of-00014.safetensors",
|
| 531 |
+
"model.language_model.layers.51.self_attn.k_proj.weight": "model-00011-of-00014.safetensors",
|
| 532 |
+
"model.language_model.layers.51.self_attn.o_proj.weight": "model-00011-of-00014.safetensors",
|
| 533 |
+
"model.language_model.layers.51.self_attn.q_norm.weight": "model-00011-of-00014.safetensors",
|
| 534 |
+
"model.language_model.layers.51.self_attn.q_proj.weight": "model-00011-of-00014.safetensors",
|
| 535 |
+
"model.language_model.layers.51.self_attn.v_proj.weight": "model-00011-of-00014.safetensors",
|
| 536 |
+
"model.language_model.layers.52.input_layernorm.weight": "model-00011-of-00014.safetensors",
|
| 537 |
+
"model.language_model.layers.52.mlp.down_proj.weight": "model-00011-of-00014.safetensors",
|
| 538 |
+
"model.language_model.layers.52.mlp.gate_proj.weight": "model-00011-of-00014.safetensors",
|
| 539 |
+
"model.language_model.layers.52.mlp.up_proj.weight": "model-00011-of-00014.safetensors",
|
| 540 |
+
"model.language_model.layers.52.post_attention_layernorm.weight": "model-00011-of-00014.safetensors",
|
| 541 |
+
"model.language_model.layers.52.self_attn.k_norm.weight": "model-00011-of-00014.safetensors",
|
| 542 |
+
"model.language_model.layers.52.self_attn.k_proj.weight": "model-00011-of-00014.safetensors",
|
| 543 |
+
"model.language_model.layers.52.self_attn.o_proj.weight": "model-00011-of-00014.safetensors",
|
| 544 |
+
"model.language_model.layers.52.self_attn.q_norm.weight": "model-00011-of-00014.safetensors",
|
| 545 |
+
"model.language_model.layers.52.self_attn.q_proj.weight": "model-00011-of-00014.safetensors",
|
| 546 |
+
"model.language_model.layers.52.self_attn.v_proj.weight": "model-00011-of-00014.safetensors",
|
| 547 |
+
"model.language_model.layers.53.input_layernorm.weight": "model-00011-of-00014.safetensors",
|
| 548 |
+
"model.language_model.layers.53.mlp.down_proj.weight": "model-00012-of-00014.safetensors",
|
| 549 |
+
"model.language_model.layers.53.mlp.gate_proj.weight": "model-00011-of-00014.safetensors",
|
| 550 |
+
"model.language_model.layers.53.mlp.up_proj.weight": "model-00012-of-00014.safetensors",
|
| 551 |
+
"model.language_model.layers.53.post_attention_layernorm.weight": "model-00011-of-00014.safetensors",
|
| 552 |
+
"model.language_model.layers.53.self_attn.k_norm.weight": "model-00011-of-00014.safetensors",
|
| 553 |
+
"model.language_model.layers.53.self_attn.k_proj.weight": "model-00011-of-00014.safetensors",
|
| 554 |
+
"model.language_model.layers.53.self_attn.o_proj.weight": "model-00011-of-00014.safetensors",
|
| 555 |
+
"model.language_model.layers.53.self_attn.q_norm.weight": "model-00011-of-00014.safetensors",
|
| 556 |
+
"model.language_model.layers.53.self_attn.q_proj.weight": "model-00011-of-00014.safetensors",
|
| 557 |
+
"model.language_model.layers.53.self_attn.v_proj.weight": "model-00011-of-00014.safetensors",
|
| 558 |
+
"model.language_model.layers.54.input_layernorm.weight": "model-00012-of-00014.safetensors",
|
| 559 |
+
"model.language_model.layers.54.mlp.down_proj.weight": "model-00012-of-00014.safetensors",
|
| 560 |
+
"model.language_model.layers.54.mlp.gate_proj.weight": "model-00012-of-00014.safetensors",
|
| 561 |
+
"model.language_model.layers.54.mlp.up_proj.weight": "model-00012-of-00014.safetensors",
|
| 562 |
+
"model.language_model.layers.54.post_attention_layernorm.weight": "model-00012-of-00014.safetensors",
|
| 563 |
+
"model.language_model.layers.54.self_attn.k_norm.weight": "model-00012-of-00014.safetensors",
|
| 564 |
+
"model.language_model.layers.54.self_attn.k_proj.weight": "model-00012-of-00014.safetensors",
|
| 565 |
+
"model.language_model.layers.54.self_attn.o_proj.weight": "model-00012-of-00014.safetensors",
|
| 566 |
+
"model.language_model.layers.54.self_attn.q_norm.weight": "model-00012-of-00014.safetensors",
|
| 567 |
+
"model.language_model.layers.54.self_attn.q_proj.weight": "model-00012-of-00014.safetensors",
|
| 568 |
+
"model.language_model.layers.54.self_attn.v_proj.weight": "model-00012-of-00014.safetensors",
|
| 569 |
+
"model.language_model.layers.55.input_layernorm.weight": "model-00012-of-00014.safetensors",
|
| 570 |
+
"model.language_model.layers.55.mlp.down_proj.weight": "model-00012-of-00014.safetensors",
|
| 571 |
+
"model.language_model.layers.55.mlp.gate_proj.weight": "model-00012-of-00014.safetensors",
|
| 572 |
+
"model.language_model.layers.55.mlp.up_proj.weight": "model-00012-of-00014.safetensors",
|
| 573 |
+
"model.language_model.layers.55.post_attention_layernorm.weight": "model-00012-of-00014.safetensors",
|
| 574 |
+
"model.language_model.layers.55.self_attn.k_norm.weight": "model-00012-of-00014.safetensors",
|
| 575 |
+
"model.language_model.layers.55.self_attn.k_proj.weight": "model-00012-of-00014.safetensors",
|
| 576 |
+
"model.language_model.layers.55.self_attn.o_proj.weight": "model-00012-of-00014.safetensors",
|
| 577 |
+
"model.language_model.layers.55.self_attn.q_norm.weight": "model-00012-of-00014.safetensors",
|
| 578 |
+
"model.language_model.layers.55.self_attn.q_proj.weight": "model-00012-of-00014.safetensors",
|
| 579 |
+
"model.language_model.layers.55.self_attn.v_proj.weight": "model-00012-of-00014.safetensors",
|
| 580 |
+
"model.language_model.layers.56.input_layernorm.weight": "model-00012-of-00014.safetensors",
|
| 581 |
+
"model.language_model.layers.56.mlp.down_proj.weight": "model-00012-of-00014.safetensors",
|
| 582 |
+
"model.language_model.layers.56.mlp.gate_proj.weight": "model-00012-of-00014.safetensors",
|
| 583 |
+
"model.language_model.layers.56.mlp.up_proj.weight": "model-00012-of-00014.safetensors",
|
| 584 |
+
"model.language_model.layers.56.post_attention_layernorm.weight": "model-00012-of-00014.safetensors",
|
| 585 |
+
"model.language_model.layers.56.self_attn.k_norm.weight": "model-00012-of-00014.safetensors",
|
| 586 |
+
"model.language_model.layers.56.self_attn.k_proj.weight": "model-00012-of-00014.safetensors",
|
| 587 |
+
"model.language_model.layers.56.self_attn.o_proj.weight": "model-00012-of-00014.safetensors",
|
| 588 |
+
"model.language_model.layers.56.self_attn.q_norm.weight": "model-00012-of-00014.safetensors",
|
| 589 |
+
"model.language_model.layers.56.self_attn.q_proj.weight": "model-00012-of-00014.safetensors",
|
| 590 |
+
"model.language_model.layers.56.self_attn.v_proj.weight": "model-00012-of-00014.safetensors",
|
| 591 |
+
"model.language_model.layers.57.input_layernorm.weight": "model-00012-of-00014.safetensors",
|
| 592 |
+
"model.language_model.layers.57.mlp.down_proj.weight": "model-00012-of-00014.safetensors",
|
| 593 |
+
"model.language_model.layers.57.mlp.gate_proj.weight": "model-00012-of-00014.safetensors",
|
| 594 |
+
"model.language_model.layers.57.mlp.up_proj.weight": "model-00012-of-00014.safetensors",
|
| 595 |
+
"model.language_model.layers.57.post_attention_layernorm.weight": "model-00012-of-00014.safetensors",
|
| 596 |
+
"model.language_model.layers.57.self_attn.k_norm.weight": "model-00012-of-00014.safetensors",
|
| 597 |
+
"model.language_model.layers.57.self_attn.k_proj.weight": "model-00012-of-00014.safetensors",
|
| 598 |
+
"model.language_model.layers.57.self_attn.o_proj.weight": "model-00012-of-00014.safetensors",
|
| 599 |
+
"model.language_model.layers.57.self_attn.q_norm.weight": "model-00012-of-00014.safetensors",
|
| 600 |
+
"model.language_model.layers.57.self_attn.q_proj.weight": "model-00012-of-00014.safetensors",
|
| 601 |
+
"model.language_model.layers.57.self_attn.v_proj.weight": "model-00012-of-00014.safetensors",
|
| 602 |
+
"model.language_model.layers.58.input_layernorm.weight": "model-00012-of-00014.safetensors",
|
| 603 |
+
"model.language_model.layers.58.mlp.down_proj.weight": "model-00013-of-00014.safetensors",
|
| 604 |
+
"model.language_model.layers.58.mlp.gate_proj.weight": "model-00012-of-00014.safetensors",
|
| 605 |
+
"model.language_model.layers.58.mlp.up_proj.weight": "model-00013-of-00014.safetensors",
|
| 606 |
+
"model.language_model.layers.58.post_attention_layernorm.weight": "model-00012-of-00014.safetensors",
|
| 607 |
+
"model.language_model.layers.58.self_attn.k_norm.weight": "model-00012-of-00014.safetensors",
|
| 608 |
+
"model.language_model.layers.58.self_attn.k_proj.weight": "model-00012-of-00014.safetensors",
|
| 609 |
+
"model.language_model.layers.58.self_attn.o_proj.weight": "model-00012-of-00014.safetensors",
|
| 610 |
+
"model.language_model.layers.58.self_attn.q_norm.weight": "model-00012-of-00014.safetensors",
|
| 611 |
+
"model.language_model.layers.58.self_attn.q_proj.weight": "model-00012-of-00014.safetensors",
|
| 612 |
+
"model.language_model.layers.58.self_attn.v_proj.weight": "model-00012-of-00014.safetensors",
|
| 613 |
+
"model.language_model.layers.59.input_layernorm.weight": "model-00013-of-00014.safetensors",
|
| 614 |
+
"model.language_model.layers.59.mlp.down_proj.weight": "model-00013-of-00014.safetensors",
|
| 615 |
+
"model.language_model.layers.59.mlp.gate_proj.weight": "model-00013-of-00014.safetensors",
|
| 616 |
+
"model.language_model.layers.59.mlp.up_proj.weight": "model-00013-of-00014.safetensors",
|
| 617 |
+
"model.language_model.layers.59.post_attention_layernorm.weight": "model-00013-of-00014.safetensors",
|
| 618 |
+
"model.language_model.layers.59.self_attn.k_norm.weight": "model-00013-of-00014.safetensors",
|
| 619 |
+
"model.language_model.layers.59.self_attn.k_proj.weight": "model-00013-of-00014.safetensors",
|
| 620 |
+
"model.language_model.layers.59.self_attn.o_proj.weight": "model-00013-of-00014.safetensors",
|
| 621 |
+
"model.language_model.layers.59.self_attn.q_norm.weight": "model-00013-of-00014.safetensors",
|
| 622 |
+
"model.language_model.layers.59.self_attn.q_proj.weight": "model-00013-of-00014.safetensors",
|
| 623 |
+
"model.language_model.layers.59.self_attn.v_proj.weight": "model-00013-of-00014.safetensors",
|
| 624 |
+
"model.language_model.layers.6.input_layernorm.weight": "model-00002-of-00014.safetensors",
|
| 625 |
+
"model.language_model.layers.6.mlp.down_proj.weight": "model-00002-of-00014.safetensors",
|
| 626 |
+
"model.language_model.layers.6.mlp.gate_proj.weight": "model-00002-of-00014.safetensors",
|
| 627 |
+
"model.language_model.layers.6.mlp.up_proj.weight": "model-00002-of-00014.safetensors",
|
| 628 |
+
"model.language_model.layers.6.post_attention_layernorm.weight": "model-00002-of-00014.safetensors",
|
| 629 |
+
"model.language_model.layers.6.self_attn.k_norm.weight": "model-00002-of-00014.safetensors",
|
| 630 |
+
"model.language_model.layers.6.self_attn.k_proj.weight": "model-00002-of-00014.safetensors",
|
| 631 |
+
"model.language_model.layers.6.self_attn.o_proj.weight": "model-00002-of-00014.safetensors",
|
| 632 |
+
"model.language_model.layers.6.self_attn.q_norm.weight": "model-00002-of-00014.safetensors",
|
| 633 |
+
"model.language_model.layers.6.self_attn.q_proj.weight": "model-00002-of-00014.safetensors",
|
| 634 |
+
"model.language_model.layers.6.self_attn.v_proj.weight": "model-00002-of-00014.safetensors",
|
| 635 |
+
"model.language_model.layers.60.input_layernorm.weight": "model-00013-of-00014.safetensors",
|
| 636 |
+
"model.language_model.layers.60.mlp.down_proj.weight": "model-00013-of-00014.safetensors",
|
| 637 |
+
"model.language_model.layers.60.mlp.gate_proj.weight": "model-00013-of-00014.safetensors",
|
| 638 |
+
"model.language_model.layers.60.mlp.up_proj.weight": "model-00013-of-00014.safetensors",
|
| 639 |
+
"model.language_model.layers.60.post_attention_layernorm.weight": "model-00013-of-00014.safetensors",
|
| 640 |
+
"model.language_model.layers.60.self_attn.k_norm.weight": "model-00013-of-00014.safetensors",
|
| 641 |
+
"model.language_model.layers.60.self_attn.k_proj.weight": "model-00013-of-00014.safetensors",
|
| 642 |
+
"model.language_model.layers.60.self_attn.o_proj.weight": "model-00013-of-00014.safetensors",
|
| 643 |
+
"model.language_model.layers.60.self_attn.q_norm.weight": "model-00013-of-00014.safetensors",
|
| 644 |
+
"model.language_model.layers.60.self_attn.q_proj.weight": "model-00013-of-00014.safetensors",
|
| 645 |
+
"model.language_model.layers.60.self_attn.v_proj.weight": "model-00013-of-00014.safetensors",
|
| 646 |
+
"model.language_model.layers.61.input_layernorm.weight": "model-00013-of-00014.safetensors",
|
| 647 |
+
"model.language_model.layers.61.mlp.down_proj.weight": "model-00013-of-00014.safetensors",
|
| 648 |
+
"model.language_model.layers.61.mlp.gate_proj.weight": "model-00013-of-00014.safetensors",
|
| 649 |
+
"model.language_model.layers.61.mlp.up_proj.weight": "model-00013-of-00014.safetensors",
|
| 650 |
+
"model.language_model.layers.61.post_attention_layernorm.weight": "model-00013-of-00014.safetensors",
|
| 651 |
+
"model.language_model.layers.61.self_attn.k_norm.weight": "model-00013-of-00014.safetensors",
|
| 652 |
+
"model.language_model.layers.61.self_attn.k_proj.weight": "model-00013-of-00014.safetensors",
|
| 653 |
+
"model.language_model.layers.61.self_attn.o_proj.weight": "model-00013-of-00014.safetensors",
|
| 654 |
+
"model.language_model.layers.61.self_attn.q_norm.weight": "model-00013-of-00014.safetensors",
|
| 655 |
+
"model.language_model.layers.61.self_attn.q_proj.weight": "model-00013-of-00014.safetensors",
|
| 656 |
+
"model.language_model.layers.61.self_attn.v_proj.weight": "model-00013-of-00014.safetensors",
|
| 657 |
+
"model.language_model.layers.62.input_layernorm.weight": "model-00013-of-00014.safetensors",
|
| 658 |
+
"model.language_model.layers.62.mlp.down_proj.weight": "model-00013-of-00014.safetensors",
|
| 659 |
+
"model.language_model.layers.62.mlp.gate_proj.weight": "model-00013-of-00014.safetensors",
|
| 660 |
+
"model.language_model.layers.62.mlp.up_proj.weight": "model-00013-of-00014.safetensors",
|
| 661 |
+
"model.language_model.layers.62.post_attention_layernorm.weight": "model-00013-of-00014.safetensors",
|
| 662 |
+
"model.language_model.layers.62.self_attn.k_norm.weight": "model-00013-of-00014.safetensors",
|
| 663 |
+
"model.language_model.layers.62.self_attn.k_proj.weight": "model-00013-of-00014.safetensors",
|
| 664 |
+
"model.language_model.layers.62.self_attn.o_proj.weight": "model-00013-of-00014.safetensors",
|
| 665 |
+
"model.language_model.layers.62.self_attn.q_norm.weight": "model-00013-of-00014.safetensors",
|
| 666 |
+
"model.language_model.layers.62.self_attn.q_proj.weight": "model-00013-of-00014.safetensors",
|
| 667 |
+
"model.language_model.layers.62.self_attn.v_proj.weight": "model-00013-of-00014.safetensors",
|
| 668 |
+
"model.language_model.layers.63.input_layernorm.weight": "model-00013-of-00014.safetensors",
|
| 669 |
+
"model.language_model.layers.63.mlp.down_proj.weight": "model-00014-of-00014.safetensors",
|
| 670 |
+
"model.language_model.layers.63.mlp.gate_proj.weight": "model-00013-of-00014.safetensors",
|
| 671 |
+
"model.language_model.layers.63.mlp.up_proj.weight": "model-00014-of-00014.safetensors",
|
| 672 |
+
"model.language_model.layers.63.post_attention_layernorm.weight": "model-00013-of-00014.safetensors",
|
| 673 |
+
"model.language_model.layers.63.self_attn.k_norm.weight": "model-00013-of-00014.safetensors",
|
| 674 |
+
"model.language_model.layers.63.self_attn.k_proj.weight": "model-00013-of-00014.safetensors",
|
| 675 |
+
"model.language_model.layers.63.self_attn.o_proj.weight": "model-00013-of-00014.safetensors",
|
| 676 |
+
"model.language_model.layers.63.self_attn.q_norm.weight": "model-00013-of-00014.safetensors",
|
| 677 |
+
"model.language_model.layers.63.self_attn.q_proj.weight": "model-00013-of-00014.safetensors",
|
| 678 |
+
"model.language_model.layers.63.self_attn.v_proj.weight": "model-00013-of-00014.safetensors",
|
| 679 |
+
"model.language_model.layers.7.input_layernorm.weight": "model-00002-of-00014.safetensors",
|
| 680 |
+
"model.language_model.layers.7.mlp.down_proj.weight": "model-00002-of-00014.safetensors",
|
| 681 |
+
"model.language_model.layers.7.mlp.gate_proj.weight": "model-00002-of-00014.safetensors",
|
| 682 |
+
"model.language_model.layers.7.mlp.up_proj.weight": "model-00002-of-00014.safetensors",
|
| 683 |
+
"model.language_model.layers.7.post_attention_layernorm.weight": "model-00002-of-00014.safetensors",
|
| 684 |
+
"model.language_model.layers.7.self_attn.k_norm.weight": "model-00002-of-00014.safetensors",
|
| 685 |
+
"model.language_model.layers.7.self_attn.k_proj.weight": "model-00002-of-00014.safetensors",
|
| 686 |
+
"model.language_model.layers.7.self_attn.o_proj.weight": "model-00002-of-00014.safetensors",
|
| 687 |
+
"model.language_model.layers.7.self_attn.q_norm.weight": "model-00002-of-00014.safetensors",
|
| 688 |
+
"model.language_model.layers.7.self_attn.q_proj.weight": "model-00002-of-00014.safetensors",
|
| 689 |
+
"model.language_model.layers.7.self_attn.v_proj.weight": "model-00002-of-00014.safetensors",
|
| 690 |
+
"model.language_model.layers.8.input_layernorm.weight": "model-00002-of-00014.safetensors",
|
| 691 |
+
"model.language_model.layers.8.mlp.down_proj.weight": "model-00003-of-00014.safetensors",
|
| 692 |
+
"model.language_model.layers.8.mlp.gate_proj.weight": "model-00002-of-00014.safetensors",
|
| 693 |
+
"model.language_model.layers.8.mlp.up_proj.weight": "model-00003-of-00014.safetensors",
|
| 694 |
+
"model.language_model.layers.8.post_attention_layernorm.weight": "model-00002-of-00014.safetensors",
|
| 695 |
+
"model.language_model.layers.8.self_attn.k_norm.weight": "model-00002-of-00014.safetensors",
|
| 696 |
+
"model.language_model.layers.8.self_attn.k_proj.weight": "model-00002-of-00014.safetensors",
|
| 697 |
+
"model.language_model.layers.8.self_attn.o_proj.weight": "model-00002-of-00014.safetensors",
|
| 698 |
+
"model.language_model.layers.8.self_attn.q_norm.weight": "model-00002-of-00014.safetensors",
|
| 699 |
+
"model.language_model.layers.8.self_attn.q_proj.weight": "model-00002-of-00014.safetensors",
|
| 700 |
+
"model.language_model.layers.8.self_attn.v_proj.weight": "model-00002-of-00014.safetensors",
|
| 701 |
+
"model.language_model.layers.9.input_layernorm.weight": "model-00003-of-00014.safetensors",
|
| 702 |
+
"model.language_model.layers.9.mlp.down_proj.weight": "model-00003-of-00014.safetensors",
|
| 703 |
+
"model.language_model.layers.9.mlp.gate_proj.weight": "model-00003-of-00014.safetensors",
|
| 704 |
+
"model.language_model.layers.9.mlp.up_proj.weight": "model-00003-of-00014.safetensors",
|
| 705 |
+
"model.language_model.layers.9.post_attention_layernorm.weight": "model-00003-of-00014.safetensors",
|
| 706 |
+
"model.language_model.layers.9.self_attn.k_norm.weight": "model-00003-of-00014.safetensors",
|
| 707 |
+
"model.language_model.layers.9.self_attn.k_proj.weight": "model-00003-of-00014.safetensors",
|
| 708 |
+
"model.language_model.layers.9.self_attn.o_proj.weight": "model-00003-of-00014.safetensors",
|
| 709 |
+
"model.language_model.layers.9.self_attn.q_norm.weight": "model-00003-of-00014.safetensors",
|
| 710 |
+
"model.language_model.layers.9.self_attn.q_proj.weight": "model-00003-of-00014.safetensors",
|
| 711 |
+
"model.language_model.layers.9.self_attn.v_proj.weight": "model-00003-of-00014.safetensors",
|
| 712 |
+
"model.language_model.norm.weight": "model-00014-of-00014.safetensors",
|
| 713 |
+
"model.visual.blocks.0.attn.proj.bias": "model-00014-of-00014.safetensors",
|
| 714 |
+
"model.visual.blocks.0.attn.proj.weight": "model-00014-of-00014.safetensors",
|
| 715 |
+
"model.visual.blocks.0.attn.qkv.bias": "model-00014-of-00014.safetensors",
|
| 716 |
+
"model.visual.blocks.0.attn.qkv.weight": "model-00014-of-00014.safetensors",
|
| 717 |
+
"model.visual.blocks.0.mlp.linear_fc1.bias": "model-00014-of-00014.safetensors",
|
| 718 |
+
"model.visual.blocks.0.mlp.linear_fc1.weight": "model-00014-of-00014.safetensors",
|
| 719 |
+
"model.visual.blocks.0.mlp.linear_fc2.bias": "model-00014-of-00014.safetensors",
|
| 720 |
+
"model.visual.blocks.0.mlp.linear_fc2.weight": "model-00014-of-00014.safetensors",
|
| 721 |
+
"model.visual.blocks.0.norm1.bias": "model-00014-of-00014.safetensors",
|
| 722 |
+
"model.visual.blocks.0.norm1.weight": "model-00014-of-00014.safetensors",
|
| 723 |
+
"model.visual.blocks.0.norm2.bias": "model-00014-of-00014.safetensors",
|
| 724 |
+
"model.visual.blocks.0.norm2.weight": "model-00014-of-00014.safetensors",
|
| 725 |
+
"model.visual.blocks.1.attn.proj.bias": "model-00014-of-00014.safetensors",
|
| 726 |
+
"model.visual.blocks.1.attn.proj.weight": "model-00014-of-00014.safetensors",
|
| 727 |
+
"model.visual.blocks.1.attn.qkv.bias": "model-00014-of-00014.safetensors",
|
| 728 |
+
"model.visual.blocks.1.attn.qkv.weight": "model-00014-of-00014.safetensors",
|
| 729 |
+
"model.visual.blocks.1.mlp.linear_fc1.bias": "model-00014-of-00014.safetensors",
|
| 730 |
+
"model.visual.blocks.1.mlp.linear_fc1.weight": "model-00014-of-00014.safetensors",
|
| 731 |
+
"model.visual.blocks.1.mlp.linear_fc2.bias": "model-00014-of-00014.safetensors",
|
| 732 |
+
"model.visual.blocks.1.mlp.linear_fc2.weight": "model-00014-of-00014.safetensors",
|
| 733 |
+
"model.visual.blocks.1.norm1.bias": "model-00014-of-00014.safetensors",
|
| 734 |
+
"model.visual.blocks.1.norm1.weight": "model-00014-of-00014.safetensors",
|
| 735 |
+
"model.visual.blocks.1.norm2.bias": "model-00014-of-00014.safetensors",
|
| 736 |
+
"model.visual.blocks.1.norm2.weight": "model-00014-of-00014.safetensors",
|
| 737 |
+
"model.visual.blocks.10.attn.proj.bias": "model-00014-of-00014.safetensors",
|
| 738 |
+
"model.visual.blocks.10.attn.proj.weight": "model-00014-of-00014.safetensors",
|
| 739 |
+
"model.visual.blocks.10.attn.qkv.bias": "model-00014-of-00014.safetensors",
|
| 740 |
+
"model.visual.blocks.10.attn.qkv.weight": "model-00014-of-00014.safetensors",
|
| 741 |
+
"model.visual.blocks.10.mlp.linear_fc1.bias": "model-00014-of-00014.safetensors",
|
| 742 |
+
"model.visual.blocks.10.mlp.linear_fc1.weight": "model-00014-of-00014.safetensors",
|
| 743 |
+
"model.visual.blocks.10.mlp.linear_fc2.bias": "model-00014-of-00014.safetensors",
|
| 744 |
+
"model.visual.blocks.10.mlp.linear_fc2.weight": "model-00014-of-00014.safetensors",
|
| 745 |
+
"model.visual.blocks.10.norm1.bias": "model-00014-of-00014.safetensors",
|
| 746 |
+
"model.visual.blocks.10.norm1.weight": "model-00014-of-00014.safetensors",
|
| 747 |
+
"model.visual.blocks.10.norm2.bias": "model-00014-of-00014.safetensors",
|
| 748 |
+
"model.visual.blocks.10.norm2.weight": "model-00014-of-00014.safetensors",
|
| 749 |
+
"model.visual.blocks.11.attn.proj.bias": "model-00014-of-00014.safetensors",
|
| 750 |
+
"model.visual.blocks.11.attn.proj.weight": "model-00014-of-00014.safetensors",
|
| 751 |
+
"model.visual.blocks.11.attn.qkv.bias": "model-00014-of-00014.safetensors",
|
| 752 |
+
"model.visual.blocks.11.attn.qkv.weight": "model-00014-of-00014.safetensors",
|
| 753 |
+
"model.visual.blocks.11.mlp.linear_fc1.bias": "model-00014-of-00014.safetensors",
|
| 754 |
+
"model.visual.blocks.11.mlp.linear_fc1.weight": "model-00014-of-00014.safetensors",
|
| 755 |
+
"model.visual.blocks.11.mlp.linear_fc2.bias": "model-00014-of-00014.safetensors",
|
| 756 |
+
"model.visual.blocks.11.mlp.linear_fc2.weight": "model-00014-of-00014.safetensors",
|
| 757 |
+
"model.visual.blocks.11.norm1.bias": "model-00014-of-00014.safetensors",
|
| 758 |
+
"model.visual.blocks.11.norm1.weight": "model-00014-of-00014.safetensors",
|
| 759 |
+
"model.visual.blocks.11.norm2.bias": "model-00014-of-00014.safetensors",
|
| 760 |
+
"model.visual.blocks.11.norm2.weight": "model-00014-of-00014.safetensors",
|
| 761 |
+
"model.visual.blocks.12.attn.proj.bias": "model-00014-of-00014.safetensors",
|
| 762 |
+
"model.visual.blocks.12.attn.proj.weight": "model-00014-of-00014.safetensors",
|
| 763 |
+
"model.visual.blocks.12.attn.qkv.bias": "model-00014-of-00014.safetensors",
|
| 764 |
+
"model.visual.blocks.12.attn.qkv.weight": "model-00014-of-00014.safetensors",
|
| 765 |
+
"model.visual.blocks.12.mlp.linear_fc1.bias": "model-00014-of-00014.safetensors",
|
| 766 |
+
"model.visual.blocks.12.mlp.linear_fc1.weight": "model-00014-of-00014.safetensors",
|
| 767 |
+
"model.visual.blocks.12.mlp.linear_fc2.bias": "model-00014-of-00014.safetensors",
|
| 768 |
+
"model.visual.blocks.12.mlp.linear_fc2.weight": "model-00014-of-00014.safetensors",
|
| 769 |
+
"model.visual.blocks.12.norm1.bias": "model-00014-of-00014.safetensors",
|
| 770 |
+
"model.visual.blocks.12.norm1.weight": "model-00014-of-00014.safetensors",
|
| 771 |
+
"model.visual.blocks.12.norm2.bias": "model-00014-of-00014.safetensors",
|
| 772 |
+
"model.visual.blocks.12.norm2.weight": "model-00014-of-00014.safetensors",
|
| 773 |
+
"model.visual.blocks.13.attn.proj.bias": "model-00014-of-00014.safetensors",
|
| 774 |
+
"model.visual.blocks.13.attn.proj.weight": "model-00014-of-00014.safetensors",
|
| 775 |
+
"model.visual.blocks.13.attn.qkv.bias": "model-00014-of-00014.safetensors",
|
| 776 |
+
"model.visual.blocks.13.attn.qkv.weight": "model-00014-of-00014.safetensors",
|
| 777 |
+
"model.visual.blocks.13.mlp.linear_fc1.bias": "model-00014-of-00014.safetensors",
|
| 778 |
+
"model.visual.blocks.13.mlp.linear_fc1.weight": "model-00014-of-00014.safetensors",
|
| 779 |
+
"model.visual.blocks.13.mlp.linear_fc2.bias": "model-00014-of-00014.safetensors",
|
| 780 |
+
"model.visual.blocks.13.mlp.linear_fc2.weight": "model-00014-of-00014.safetensors",
|
| 781 |
+
"model.visual.blocks.13.norm1.bias": "model-00014-of-00014.safetensors",
|
| 782 |
+
"model.visual.blocks.13.norm1.weight": "model-00014-of-00014.safetensors",
|
| 783 |
+
"model.visual.blocks.13.norm2.bias": "model-00014-of-00014.safetensors",
|
| 784 |
+
"model.visual.blocks.13.norm2.weight": "model-00014-of-00014.safetensors",
|
| 785 |
+
"model.visual.blocks.14.attn.proj.bias": "model-00014-of-00014.safetensors",
|
| 786 |
+
"model.visual.blocks.14.attn.proj.weight": "model-00014-of-00014.safetensors",
|
| 787 |
+
"model.visual.blocks.14.attn.qkv.bias": "model-00014-of-00014.safetensors",
|
| 788 |
+
"model.visual.blocks.14.attn.qkv.weight": "model-00014-of-00014.safetensors",
|
| 789 |
+
"model.visual.blocks.14.mlp.linear_fc1.bias": "model-00014-of-00014.safetensors",
|
| 790 |
+
"model.visual.blocks.14.mlp.linear_fc1.weight": "model-00014-of-00014.safetensors",
|
| 791 |
+
"model.visual.blocks.14.mlp.linear_fc2.bias": "model-00014-of-00014.safetensors",
|
| 792 |
+
"model.visual.blocks.14.mlp.linear_fc2.weight": "model-00014-of-00014.safetensors",
|
| 793 |
+
"model.visual.blocks.14.norm1.bias": "model-00014-of-00014.safetensors",
|
| 794 |
+
"model.visual.blocks.14.norm1.weight": "model-00014-of-00014.safetensors",
|
| 795 |
+
"model.visual.blocks.14.norm2.bias": "model-00014-of-00014.safetensors",
|
| 796 |
+
"model.visual.blocks.14.norm2.weight": "model-00014-of-00014.safetensors",
|
| 797 |
+
"model.visual.blocks.15.attn.proj.bias": "model-00014-of-00014.safetensors",
|
| 798 |
+
"model.visual.blocks.15.attn.proj.weight": "model-00014-of-00014.safetensors",
|
| 799 |
+
"model.visual.blocks.15.attn.qkv.bias": "model-00014-of-00014.safetensors",
|
| 800 |
+
"model.visual.blocks.15.attn.qkv.weight": "model-00014-of-00014.safetensors",
|
| 801 |
+
"model.visual.blocks.15.mlp.linear_fc1.bias": "model-00014-of-00014.safetensors",
|
| 802 |
+
"model.visual.blocks.15.mlp.linear_fc1.weight": "model-00014-of-00014.safetensors",
|
| 803 |
+
"model.visual.blocks.15.mlp.linear_fc2.bias": "model-00014-of-00014.safetensors",
|
| 804 |
+
"model.visual.blocks.15.mlp.linear_fc2.weight": "model-00014-of-00014.safetensors",
|
| 805 |
+
"model.visual.blocks.15.norm1.bias": "model-00014-of-00014.safetensors",
|
| 806 |
+
"model.visual.blocks.15.norm1.weight": "model-00014-of-00014.safetensors",
|
| 807 |
+
"model.visual.blocks.15.norm2.bias": "model-00014-of-00014.safetensors",
|
| 808 |
+
"model.visual.blocks.15.norm2.weight": "model-00014-of-00014.safetensors",
|
| 809 |
+
"model.visual.blocks.16.attn.proj.bias": "model-00014-of-00014.safetensors",
|
| 810 |
+
"model.visual.blocks.16.attn.proj.weight": "model-00014-of-00014.safetensors",
|
| 811 |
+
"model.visual.blocks.16.attn.qkv.bias": "model-00014-of-00014.safetensors",
|
| 812 |
+
"model.visual.blocks.16.attn.qkv.weight": "model-00014-of-00014.safetensors",
|
| 813 |
+
"model.visual.blocks.16.mlp.linear_fc1.bias": "model-00014-of-00014.safetensors",
|
| 814 |
+
"model.visual.blocks.16.mlp.linear_fc1.weight": "model-00014-of-00014.safetensors",
|
| 815 |
+
"model.visual.blocks.16.mlp.linear_fc2.bias": "model-00014-of-00014.safetensors",
|
| 816 |
+
"model.visual.blocks.16.mlp.linear_fc2.weight": "model-00014-of-00014.safetensors",
|
| 817 |
+
"model.visual.blocks.16.norm1.bias": "model-00014-of-00014.safetensors",
|
| 818 |
+
"model.visual.blocks.16.norm1.weight": "model-00014-of-00014.safetensors",
|
| 819 |
+
"model.visual.blocks.16.norm2.bias": "model-00014-of-00014.safetensors",
|
| 820 |
+
"model.visual.blocks.16.norm2.weight": "model-00014-of-00014.safetensors",
|
| 821 |
+
"model.visual.blocks.17.attn.proj.bias": "model-00014-of-00014.safetensors",
|
| 822 |
+
"model.visual.blocks.17.attn.proj.weight": "model-00014-of-00014.safetensors",
|
| 823 |
+
"model.visual.blocks.17.attn.qkv.bias": "model-00014-of-00014.safetensors",
|
| 824 |
+
"model.visual.blocks.17.attn.qkv.weight": "model-00014-of-00014.safetensors",
|
| 825 |
+
"model.visual.blocks.17.mlp.linear_fc1.bias": "model-00014-of-00014.safetensors",
|
| 826 |
+
"model.visual.blocks.17.mlp.linear_fc1.weight": "model-00014-of-00014.safetensors",
|
| 827 |
+
"model.visual.blocks.17.mlp.linear_fc2.bias": "model-00014-of-00014.safetensors",
|
| 828 |
+
"model.visual.blocks.17.mlp.linear_fc2.weight": "model-00014-of-00014.safetensors",
|
| 829 |
+
"model.visual.blocks.17.norm1.bias": "model-00014-of-00014.safetensors",
|
| 830 |
+
"model.visual.blocks.17.norm1.weight": "model-00014-of-00014.safetensors",
|
| 831 |
+
"model.visual.blocks.17.norm2.bias": "model-00014-of-00014.safetensors",
|
| 832 |
+
"model.visual.blocks.17.norm2.weight": "model-00014-of-00014.safetensors",
|
| 833 |
+
"model.visual.blocks.18.attn.proj.bias": "model-00014-of-00014.safetensors",
|
| 834 |
+
"model.visual.blocks.18.attn.proj.weight": "model-00014-of-00014.safetensors",
|
| 835 |
+
"model.visual.blocks.18.attn.qkv.bias": "model-00014-of-00014.safetensors",
|
| 836 |
+
"model.visual.blocks.18.attn.qkv.weight": "model-00014-of-00014.safetensors",
|
| 837 |
+
"model.visual.blocks.18.mlp.linear_fc1.bias": "model-00014-of-00014.safetensors",
|
| 838 |
+
"model.visual.blocks.18.mlp.linear_fc1.weight": "model-00014-of-00014.safetensors",
|
| 839 |
+
"model.visual.blocks.18.mlp.linear_fc2.bias": "model-00014-of-00014.safetensors",
|
| 840 |
+
"model.visual.blocks.18.mlp.linear_fc2.weight": "model-00014-of-00014.safetensors",
|
| 841 |
+
"model.visual.blocks.18.norm1.bias": "model-00014-of-00014.safetensors",
|
| 842 |
+
"model.visual.blocks.18.norm1.weight": "model-00014-of-00014.safetensors",
|
| 843 |
+
"model.visual.blocks.18.norm2.bias": "model-00014-of-00014.safetensors",
|
| 844 |
+
"model.visual.blocks.18.norm2.weight": "model-00014-of-00014.safetensors",
|
| 845 |
+
"model.visual.blocks.19.attn.proj.bias": "model-00014-of-00014.safetensors",
|
| 846 |
+
"model.visual.blocks.19.attn.proj.weight": "model-00014-of-00014.safetensors",
|
| 847 |
+
"model.visual.blocks.19.attn.qkv.bias": "model-00014-of-00014.safetensors",
|
| 848 |
+
"model.visual.blocks.19.attn.qkv.weight": "model-00014-of-00014.safetensors",
|
| 849 |
+
"model.visual.blocks.19.mlp.linear_fc1.bias": "model-00014-of-00014.safetensors",
|
| 850 |
+
"model.visual.blocks.19.mlp.linear_fc1.weight": "model-00014-of-00014.safetensors",
|
| 851 |
+
"model.visual.blocks.19.mlp.linear_fc2.bias": "model-00014-of-00014.safetensors",
|
| 852 |
+
"model.visual.blocks.19.mlp.linear_fc2.weight": "model-00014-of-00014.safetensors",
|
| 853 |
+
"model.visual.blocks.19.norm1.bias": "model-00014-of-00014.safetensors",
|
| 854 |
+
"model.visual.blocks.19.norm1.weight": "model-00014-of-00014.safetensors",
|
| 855 |
+
"model.visual.blocks.19.norm2.bias": "model-00014-of-00014.safetensors",
|
| 856 |
+
"model.visual.blocks.19.norm2.weight": "model-00014-of-00014.safetensors",
|
| 857 |
+
"model.visual.blocks.2.attn.proj.bias": "model-00014-of-00014.safetensors",
|
| 858 |
+
"model.visual.blocks.2.attn.proj.weight": "model-00014-of-00014.safetensors",
|
| 859 |
+
"model.visual.blocks.2.attn.qkv.bias": "model-00014-of-00014.safetensors",
|
| 860 |
+
"model.visual.blocks.2.attn.qkv.weight": "model-00014-of-00014.safetensors",
|
| 861 |
+
"model.visual.blocks.2.mlp.linear_fc1.bias": "model-00014-of-00014.safetensors",
|
| 862 |
+
"model.visual.blocks.2.mlp.linear_fc1.weight": "model-00014-of-00014.safetensors",
|
| 863 |
+
"model.visual.blocks.2.mlp.linear_fc2.bias": "model-00014-of-00014.safetensors",
|
| 864 |
+
"model.visual.blocks.2.mlp.linear_fc2.weight": "model-00014-of-00014.safetensors",
|
| 865 |
+
"model.visual.blocks.2.norm1.bias": "model-00014-of-00014.safetensors",
|
| 866 |
+
"model.visual.blocks.2.norm1.weight": "model-00014-of-00014.safetensors",
|
| 867 |
+
"model.visual.blocks.2.norm2.bias": "model-00014-of-00014.safetensors",
|
| 868 |
+
"model.visual.blocks.2.norm2.weight": "model-00014-of-00014.safetensors",
|
| 869 |
+
"model.visual.blocks.20.attn.proj.bias": "model-00014-of-00014.safetensors",
|
| 870 |
+
"model.visual.blocks.20.attn.proj.weight": "model-00014-of-00014.safetensors",
|
| 871 |
+
"model.visual.blocks.20.attn.qkv.bias": "model-00014-of-00014.safetensors",
|
| 872 |
+
"model.visual.blocks.20.attn.qkv.weight": "model-00014-of-00014.safetensors",
|
| 873 |
+
"model.visual.blocks.20.mlp.linear_fc1.bias": "model-00014-of-00014.safetensors",
|
| 874 |
+
"model.visual.blocks.20.mlp.linear_fc1.weight": "model-00014-of-00014.safetensors",
|
| 875 |
+
"model.visual.blocks.20.mlp.linear_fc2.bias": "model-00014-of-00014.safetensors",
|
| 876 |
+
"model.visual.blocks.20.mlp.linear_fc2.weight": "model-00014-of-00014.safetensors",
|
| 877 |
+
"model.visual.blocks.20.norm1.bias": "model-00014-of-00014.safetensors",
|
| 878 |
+
"model.visual.blocks.20.norm1.weight": "model-00014-of-00014.safetensors",
|
| 879 |
+
"model.visual.blocks.20.norm2.bias": "model-00014-of-00014.safetensors",
|
| 880 |
+
"model.visual.blocks.20.norm2.weight": "model-00014-of-00014.safetensors",
|
| 881 |
+
"model.visual.blocks.21.attn.proj.bias": "model-00014-of-00014.safetensors",
|
| 882 |
+
"model.visual.blocks.21.attn.proj.weight": "model-00014-of-00014.safetensors",
|
| 883 |
+
"model.visual.blocks.21.attn.qkv.bias": "model-00014-of-00014.safetensors",
|
| 884 |
+
"model.visual.blocks.21.attn.qkv.weight": "model-00014-of-00014.safetensors",
|
| 885 |
+
"model.visual.blocks.21.mlp.linear_fc1.bias": "model-00014-of-00014.safetensors",
|
| 886 |
+
"model.visual.blocks.21.mlp.linear_fc1.weight": "model-00014-of-00014.safetensors",
|
| 887 |
+
"model.visual.blocks.21.mlp.linear_fc2.bias": "model-00014-of-00014.safetensors",
|
| 888 |
+
"model.visual.blocks.21.mlp.linear_fc2.weight": "model-00014-of-00014.safetensors",
|
| 889 |
+
"model.visual.blocks.21.norm1.bias": "model-00014-of-00014.safetensors",
|
| 890 |
+
"model.visual.blocks.21.norm1.weight": "model-00014-of-00014.safetensors",
|
| 891 |
+
"model.visual.blocks.21.norm2.bias": "model-00014-of-00014.safetensors",
|
| 892 |
+
"model.visual.blocks.21.norm2.weight": "model-00014-of-00014.safetensors",
|
| 893 |
+
"model.visual.blocks.22.attn.proj.bias": "model-00014-of-00014.safetensors",
|
| 894 |
+
"model.visual.blocks.22.attn.proj.weight": "model-00014-of-00014.safetensors",
|
| 895 |
+
"model.visual.blocks.22.attn.qkv.bias": "model-00014-of-00014.safetensors",
|
| 896 |
+
"model.visual.blocks.22.attn.qkv.weight": "model-00014-of-00014.safetensors",
|
| 897 |
+
"model.visual.blocks.22.mlp.linear_fc1.bias": "model-00014-of-00014.safetensors",
|
| 898 |
+
"model.visual.blocks.22.mlp.linear_fc1.weight": "model-00014-of-00014.safetensors",
|
| 899 |
+
"model.visual.blocks.22.mlp.linear_fc2.bias": "model-00014-of-00014.safetensors",
|
| 900 |
+
"model.visual.blocks.22.mlp.linear_fc2.weight": "model-00014-of-00014.safetensors",
|
| 901 |
+
"model.visual.blocks.22.norm1.bias": "model-00014-of-00014.safetensors",
|
| 902 |
+
"model.visual.blocks.22.norm1.weight": "model-00014-of-00014.safetensors",
|
| 903 |
+
"model.visual.blocks.22.norm2.bias": "model-00014-of-00014.safetensors",
|
| 904 |
+
"model.visual.blocks.22.norm2.weight": "model-00014-of-00014.safetensors",
|
| 905 |
+
"model.visual.blocks.23.attn.proj.bias": "model-00014-of-00014.safetensors",
|
| 906 |
+
"model.visual.blocks.23.attn.proj.weight": "model-00014-of-00014.safetensors",
|
| 907 |
+
"model.visual.blocks.23.attn.qkv.bias": "model-00014-of-00014.safetensors",
|
| 908 |
+
"model.visual.blocks.23.attn.qkv.weight": "model-00014-of-00014.safetensors",
|
| 909 |
+
"model.visual.blocks.23.mlp.linear_fc1.bias": "model-00014-of-00014.safetensors",
|
| 910 |
+
"model.visual.blocks.23.mlp.linear_fc1.weight": "model-00014-of-00014.safetensors",
|
| 911 |
+
"model.visual.blocks.23.mlp.linear_fc2.bias": "model-00014-of-00014.safetensors",
|
| 912 |
+
"model.visual.blocks.23.mlp.linear_fc2.weight": "model-00014-of-00014.safetensors",
|
| 913 |
+
"model.visual.blocks.23.norm1.bias": "model-00014-of-00014.safetensors",
|
| 914 |
+
"model.visual.blocks.23.norm1.weight": "model-00014-of-00014.safetensors",
|
| 915 |
+
"model.visual.blocks.23.norm2.bias": "model-00014-of-00014.safetensors",
|
| 916 |
+
"model.visual.blocks.23.norm2.weight": "model-00014-of-00014.safetensors",
|
| 917 |
+
"model.visual.blocks.24.attn.proj.bias": "model-00014-of-00014.safetensors",
|
| 918 |
+
"model.visual.blocks.24.attn.proj.weight": "model-00014-of-00014.safetensors",
|
| 919 |
+
"model.visual.blocks.24.attn.qkv.bias": "model-00014-of-00014.safetensors",
|
| 920 |
+
"model.visual.blocks.24.attn.qkv.weight": "model-00014-of-00014.safetensors",
|
| 921 |
+
"model.visual.blocks.24.mlp.linear_fc1.bias": "model-00014-of-00014.safetensors",
|
| 922 |
+
"model.visual.blocks.24.mlp.linear_fc1.weight": "model-00014-of-00014.safetensors",
|
| 923 |
+
"model.visual.blocks.24.mlp.linear_fc2.bias": "model-00014-of-00014.safetensors",
|
| 924 |
+
"model.visual.blocks.24.mlp.linear_fc2.weight": "model-00014-of-00014.safetensors",
|
| 925 |
+
"model.visual.blocks.24.norm1.bias": "model-00014-of-00014.safetensors",
|
| 926 |
+
"model.visual.blocks.24.norm1.weight": "model-00014-of-00014.safetensors",
|
| 927 |
+
"model.visual.blocks.24.norm2.bias": "model-00014-of-00014.safetensors",
|
| 928 |
+
"model.visual.blocks.24.norm2.weight": "model-00014-of-00014.safetensors",
|
| 929 |
+
"model.visual.blocks.25.attn.proj.bias": "model-00014-of-00014.safetensors",
|
| 930 |
+
"model.visual.blocks.25.attn.proj.weight": "model-00014-of-00014.safetensors",
|
| 931 |
+
"model.visual.blocks.25.attn.qkv.bias": "model-00014-of-00014.safetensors",
|
| 932 |
+
"model.visual.blocks.25.attn.qkv.weight": "model-00014-of-00014.safetensors",
|
| 933 |
+
"model.visual.blocks.25.mlp.linear_fc1.bias": "model-00014-of-00014.safetensors",
|
| 934 |
+
"model.visual.blocks.25.mlp.linear_fc1.weight": "model-00014-of-00014.safetensors",
|
| 935 |
+
"model.visual.blocks.25.mlp.linear_fc2.bias": "model-00014-of-00014.safetensors",
|
| 936 |
+
"model.visual.blocks.25.mlp.linear_fc2.weight": "model-00014-of-00014.safetensors",
|
| 937 |
+
"model.visual.blocks.25.norm1.bias": "model-00014-of-00014.safetensors",
|
| 938 |
+
"model.visual.blocks.25.norm1.weight": "model-00014-of-00014.safetensors",
|
| 939 |
+
"model.visual.blocks.25.norm2.bias": "model-00014-of-00014.safetensors",
|
| 940 |
+
"model.visual.blocks.25.norm2.weight": "model-00014-of-00014.safetensors",
|
| 941 |
+
"model.visual.blocks.26.attn.proj.bias": "model-00014-of-00014.safetensors",
|
| 942 |
+
"model.visual.blocks.26.attn.proj.weight": "model-00014-of-00014.safetensors",
|
| 943 |
+
"model.visual.blocks.26.attn.qkv.bias": "model-00014-of-00014.safetensors",
|
| 944 |
+
"model.visual.blocks.26.attn.qkv.weight": "model-00014-of-00014.safetensors",
|
| 945 |
+
"model.visual.blocks.26.mlp.linear_fc1.bias": "model-00014-of-00014.safetensors",
|
| 946 |
+
"model.visual.blocks.26.mlp.linear_fc1.weight": "model-00014-of-00014.safetensors",
|
| 947 |
+
"model.visual.blocks.26.mlp.linear_fc2.bias": "model-00014-of-00014.safetensors",
|
| 948 |
+
"model.visual.blocks.26.mlp.linear_fc2.weight": "model-00014-of-00014.safetensors",
|
| 949 |
+
"model.visual.blocks.26.norm1.bias": "model-00014-of-00014.safetensors",
|
| 950 |
+
"model.visual.blocks.26.norm1.weight": "model-00014-of-00014.safetensors",
|
| 951 |
+
"model.visual.blocks.26.norm2.bias": "model-00014-of-00014.safetensors",
|
| 952 |
+
"model.visual.blocks.26.norm2.weight": "model-00014-of-00014.safetensors",
|
| 953 |
+
"model.visual.blocks.3.attn.proj.bias": "model-00014-of-00014.safetensors",
|
| 954 |
+
"model.visual.blocks.3.attn.proj.weight": "model-00014-of-00014.safetensors",
|
| 955 |
+
"model.visual.blocks.3.attn.qkv.bias": "model-00014-of-00014.safetensors",
|
| 956 |
+
"model.visual.blocks.3.attn.qkv.weight": "model-00014-of-00014.safetensors",
|
| 957 |
+
"model.visual.blocks.3.mlp.linear_fc1.bias": "model-00014-of-00014.safetensors",
|
| 958 |
+
"model.visual.blocks.3.mlp.linear_fc1.weight": "model-00014-of-00014.safetensors",
|
| 959 |
+
"model.visual.blocks.3.mlp.linear_fc2.bias": "model-00014-of-00014.safetensors",
|
| 960 |
+
"model.visual.blocks.3.mlp.linear_fc2.weight": "model-00014-of-00014.safetensors",
|
| 961 |
+
"model.visual.blocks.3.norm1.bias": "model-00014-of-00014.safetensors",
|
| 962 |
+
"model.visual.blocks.3.norm1.weight": "model-00014-of-00014.safetensors",
|
| 963 |
+
"model.visual.blocks.3.norm2.bias": "model-00014-of-00014.safetensors",
|
| 964 |
+
"model.visual.blocks.3.norm2.weight": "model-00014-of-00014.safetensors",
|
| 965 |
+
"model.visual.blocks.4.attn.proj.bias": "model-00014-of-00014.safetensors",
|
| 966 |
+
"model.visual.blocks.4.attn.proj.weight": "model-00014-of-00014.safetensors",
|
| 967 |
+
"model.visual.blocks.4.attn.qkv.bias": "model-00014-of-00014.safetensors",
|
| 968 |
+
"model.visual.blocks.4.attn.qkv.weight": "model-00014-of-00014.safetensors",
|
| 969 |
+
"model.visual.blocks.4.mlp.linear_fc1.bias": "model-00014-of-00014.safetensors",
|
| 970 |
+
"model.visual.blocks.4.mlp.linear_fc1.weight": "model-00014-of-00014.safetensors",
|
| 971 |
+
"model.visual.blocks.4.mlp.linear_fc2.bias": "model-00014-of-00014.safetensors",
|
| 972 |
+
"model.visual.blocks.4.mlp.linear_fc2.weight": "model-00014-of-00014.safetensors",
|
| 973 |
+
"model.visual.blocks.4.norm1.bias": "model-00014-of-00014.safetensors",
|
| 974 |
+
"model.visual.blocks.4.norm1.weight": "model-00014-of-00014.safetensors",
|
| 975 |
+
"model.visual.blocks.4.norm2.bias": "model-00014-of-00014.safetensors",
|
| 976 |
+
"model.visual.blocks.4.norm2.weight": "model-00014-of-00014.safetensors",
|
| 977 |
+
"model.visual.blocks.5.attn.proj.bias": "model-00014-of-00014.safetensors",
|
| 978 |
+
"model.visual.blocks.5.attn.proj.weight": "model-00014-of-00014.safetensors",
|
| 979 |
+
"model.visual.blocks.5.attn.qkv.bias": "model-00014-of-00014.safetensors",
|
| 980 |
+
"model.visual.blocks.5.attn.qkv.weight": "model-00014-of-00014.safetensors",
|
| 981 |
+
"model.visual.blocks.5.mlp.linear_fc1.bias": "model-00014-of-00014.safetensors",
|
| 982 |
+
"model.visual.blocks.5.mlp.linear_fc1.weight": "model-00014-of-00014.safetensors",
|
| 983 |
+
"model.visual.blocks.5.mlp.linear_fc2.bias": "model-00014-of-00014.safetensors",
|
| 984 |
+
"model.visual.blocks.5.mlp.linear_fc2.weight": "model-00014-of-00014.safetensors",
|
| 985 |
+
"model.visual.blocks.5.norm1.bias": "model-00014-of-00014.safetensors",
|
| 986 |
+
"model.visual.blocks.5.norm1.weight": "model-00014-of-00014.safetensors",
|
| 987 |
+
"model.visual.blocks.5.norm2.bias": "model-00014-of-00014.safetensors",
|
| 988 |
+
"model.visual.blocks.5.norm2.weight": "model-00014-of-00014.safetensors",
|
| 989 |
+
"model.visual.blocks.6.attn.proj.bias": "model-00014-of-00014.safetensors",
|
| 990 |
+
"model.visual.blocks.6.attn.proj.weight": "model-00014-of-00014.safetensors",
|
| 991 |
+
"model.visual.blocks.6.attn.qkv.bias": "model-00014-of-00014.safetensors",
|
| 992 |
+
"model.visual.blocks.6.attn.qkv.weight": "model-00014-of-00014.safetensors",
|
| 993 |
+
"model.visual.blocks.6.mlp.linear_fc1.bias": "model-00014-of-00014.safetensors",
|
| 994 |
+
"model.visual.blocks.6.mlp.linear_fc1.weight": "model-00014-of-00014.safetensors",
|
| 995 |
+
"model.visual.blocks.6.mlp.linear_fc2.bias": "model-00014-of-00014.safetensors",
|
| 996 |
+
"model.visual.blocks.6.mlp.linear_fc2.weight": "model-00014-of-00014.safetensors",
|
| 997 |
+
"model.visual.blocks.6.norm1.bias": "model-00014-of-00014.safetensors",
|
| 998 |
+
"model.visual.blocks.6.norm1.weight": "model-00014-of-00014.safetensors",
|
| 999 |
+
"model.visual.blocks.6.norm2.bias": "model-00014-of-00014.safetensors",
|
| 1000 |
+
"model.visual.blocks.6.norm2.weight": "model-00014-of-00014.safetensors",
|
| 1001 |
+
"model.visual.blocks.7.attn.proj.bias": "model-00014-of-00014.safetensors",
|
| 1002 |
+
"model.visual.blocks.7.attn.proj.weight": "model-00014-of-00014.safetensors",
|
| 1003 |
+
"model.visual.blocks.7.attn.qkv.bias": "model-00014-of-00014.safetensors",
|
| 1004 |
+
"model.visual.blocks.7.attn.qkv.weight": "model-00014-of-00014.safetensors",
|
| 1005 |
+
"model.visual.blocks.7.mlp.linear_fc1.bias": "model-00014-of-00014.safetensors",
|
| 1006 |
+
"model.visual.blocks.7.mlp.linear_fc1.weight": "model-00014-of-00014.safetensors",
|
| 1007 |
+
"model.visual.blocks.7.mlp.linear_fc2.bias": "model-00014-of-00014.safetensors",
|
| 1008 |
+
"model.visual.blocks.7.mlp.linear_fc2.weight": "model-00014-of-00014.safetensors",
|
| 1009 |
+
"model.visual.blocks.7.norm1.bias": "model-00014-of-00014.safetensors",
|
| 1010 |
+
"model.visual.blocks.7.norm1.weight": "model-00014-of-00014.safetensors",
|
| 1011 |
+
"model.visual.blocks.7.norm2.bias": "model-00014-of-00014.safetensors",
|
| 1012 |
+
"model.visual.blocks.7.norm2.weight": "model-00014-of-00014.safetensors",
|
| 1013 |
+
"model.visual.blocks.8.attn.proj.bias": "model-00014-of-00014.safetensors",
|
| 1014 |
+
"model.visual.blocks.8.attn.proj.weight": "model-00014-of-00014.safetensors",
|
| 1015 |
+
"model.visual.blocks.8.attn.qkv.bias": "model-00014-of-00014.safetensors",
|
| 1016 |
+
"model.visual.blocks.8.attn.qkv.weight": "model-00014-of-00014.safetensors",
|
| 1017 |
+
"model.visual.blocks.8.mlp.linear_fc1.bias": "model-00014-of-00014.safetensors",
|
| 1018 |
+
"model.visual.blocks.8.mlp.linear_fc1.weight": "model-00014-of-00014.safetensors",
|
| 1019 |
+
"model.visual.blocks.8.mlp.linear_fc2.bias": "model-00014-of-00014.safetensors",
|
| 1020 |
+
"model.visual.blocks.8.mlp.linear_fc2.weight": "model-00014-of-00014.safetensors",
|
| 1021 |
+
"model.visual.blocks.8.norm1.bias": "model-00014-of-00014.safetensors",
|
| 1022 |
+
"model.visual.blocks.8.norm1.weight": "model-00014-of-00014.safetensors",
|
| 1023 |
+
"model.visual.blocks.8.norm2.bias": "model-00014-of-00014.safetensors",
|
| 1024 |
+
"model.visual.blocks.8.norm2.weight": "model-00014-of-00014.safetensors",
|
| 1025 |
+
"model.visual.blocks.9.attn.proj.bias": "model-00014-of-00014.safetensors",
|
| 1026 |
+
"model.visual.blocks.9.attn.proj.weight": "model-00014-of-00014.safetensors",
|
| 1027 |
+
"model.visual.blocks.9.attn.qkv.bias": "model-00014-of-00014.safetensors",
|
| 1028 |
+
"model.visual.blocks.9.attn.qkv.weight": "model-00014-of-00014.safetensors",
|
| 1029 |
+
"model.visual.blocks.9.mlp.linear_fc1.bias": "model-00014-of-00014.safetensors",
|
| 1030 |
+
"model.visual.blocks.9.mlp.linear_fc1.weight": "model-00014-of-00014.safetensors",
|
| 1031 |
+
"model.visual.blocks.9.mlp.linear_fc2.bias": "model-00014-of-00014.safetensors",
|
| 1032 |
+
"model.visual.blocks.9.mlp.linear_fc2.weight": "model-00014-of-00014.safetensors",
|
| 1033 |
+
"model.visual.blocks.9.norm1.bias": "model-00014-of-00014.safetensors",
|
| 1034 |
+
"model.visual.blocks.9.norm1.weight": "model-00014-of-00014.safetensors",
|
| 1035 |
+
"model.visual.blocks.9.norm2.bias": "model-00014-of-00014.safetensors",
|
| 1036 |
+
"model.visual.blocks.9.norm2.weight": "model-00014-of-00014.safetensors",
|
| 1037 |
+
"model.visual.deepstack_merger_list.0.linear_fc1.bias": "model-00014-of-00014.safetensors",
|
| 1038 |
+
"model.visual.deepstack_merger_list.0.linear_fc1.weight": "model-00014-of-00014.safetensors",
|
| 1039 |
+
"model.visual.deepstack_merger_list.0.linear_fc2.bias": "model-00014-of-00014.safetensors",
|
| 1040 |
+
"model.visual.deepstack_merger_list.0.linear_fc2.weight": "model-00014-of-00014.safetensors",
|
| 1041 |
+
"model.visual.deepstack_merger_list.0.norm.bias": "model-00014-of-00014.safetensors",
|
| 1042 |
+
"model.visual.deepstack_merger_list.0.norm.weight": "model-00014-of-00014.safetensors",
|
| 1043 |
+
"model.visual.deepstack_merger_list.1.linear_fc1.bias": "model-00014-of-00014.safetensors",
|
| 1044 |
+
"model.visual.deepstack_merger_list.1.linear_fc1.weight": "model-00014-of-00014.safetensors",
|
| 1045 |
+
"model.visual.deepstack_merger_list.1.linear_fc2.bias": "model-00014-of-00014.safetensors",
|
| 1046 |
+
"model.visual.deepstack_merger_list.1.linear_fc2.weight": "model-00014-of-00014.safetensors",
|
| 1047 |
+
"model.visual.deepstack_merger_list.1.norm.bias": "model-00014-of-00014.safetensors",
|
| 1048 |
+
"model.visual.deepstack_merger_list.1.norm.weight": "model-00014-of-00014.safetensors",
|
| 1049 |
+
"model.visual.deepstack_merger_list.2.linear_fc1.bias": "model-00014-of-00014.safetensors",
|
| 1050 |
+
"model.visual.deepstack_merger_list.2.linear_fc1.weight": "model-00014-of-00014.safetensors",
|
| 1051 |
+
"model.visual.deepstack_merger_list.2.linear_fc2.bias": "model-00014-of-00014.safetensors",
|
| 1052 |
+
"model.visual.deepstack_merger_list.2.linear_fc2.weight": "model-00014-of-00014.safetensors",
|
| 1053 |
+
"model.visual.deepstack_merger_list.2.norm.bias": "model-00014-of-00014.safetensors",
|
| 1054 |
+
"model.visual.deepstack_merger_list.2.norm.weight": "model-00014-of-00014.safetensors",
|
| 1055 |
+
"model.visual.merger.linear_fc1.bias": "model-00014-of-00014.safetensors",
|
| 1056 |
+
"model.visual.merger.linear_fc1.weight": "model-00014-of-00014.safetensors",
|
| 1057 |
+
"model.visual.merger.linear_fc2.bias": "model-00014-of-00014.safetensors",
|
| 1058 |
+
"model.visual.merger.linear_fc2.weight": "model-00014-of-00014.safetensors",
|
| 1059 |
+
"model.visual.merger.norm.bias": "model-00014-of-00014.safetensors",
|
| 1060 |
+
"model.visual.merger.norm.weight": "model-00014-of-00014.safetensors",
|
| 1061 |
+
"model.visual.patch_embed.proj.bias": "model-00014-of-00014.safetensors",
|
| 1062 |
+
"model.visual.patch_embed.proj.weight": "model-00014-of-00014.safetensors",
|
| 1063 |
+
"model.visual.pos_embed.weight": "model-00014-of-00014.safetensors"
|
| 1064 |
+
}
|
| 1065 |
+
}
|
FL2VA/text_encoder/preprocessor_config.json
ADDED
|
@@ -0,0 +1,21 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"size": {
|
| 3 |
+
"longest_edge": 16777216,
|
| 4 |
+
"shortest_edge": 65536
|
| 5 |
+
},
|
| 6 |
+
"patch_size": 16,
|
| 7 |
+
"temporal_patch_size": 2,
|
| 8 |
+
"merge_size": 2,
|
| 9 |
+
"image_mean": [
|
| 10 |
+
0.5,
|
| 11 |
+
0.5,
|
| 12 |
+
0.5
|
| 13 |
+
],
|
| 14 |
+
"image_std": [
|
| 15 |
+
0.5,
|
| 16 |
+
0.5,
|
| 17 |
+
0.5
|
| 18 |
+
],
|
| 19 |
+
"processor_class": "Qwen3VLProcessor",
|
| 20 |
+
"image_processor_type": "Qwen2VLImageProcessorFast"
|
| 21 |
+
}
|
FL2VA/text_encoder/tokenizer.json
ADDED
|
The diff for this file is too large to render.
See raw diff
|
|
|
FL2VA/text_encoder/tokenizer_config.json
ADDED
|
@@ -0,0 +1,246 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"add_bos_token": false,
|
| 3 |
+
"add_prefix_space": false,
|
| 4 |
+
"added_tokens_decoder": {
|
| 5 |
+
"151643": {
|
| 6 |
+
"content": "<|endoftext|>",
|
| 7 |
+
"lstrip": false,
|
| 8 |
+
"normalized": false,
|
| 9 |
+
"rstrip": false,
|
| 10 |
+
"single_word": false,
|
| 11 |
+
"special": true
|
| 12 |
+
},
|
| 13 |
+
"151644": {
|
| 14 |
+
"content": "<|im_start|>",
|
| 15 |
+
"lstrip": false,
|
| 16 |
+
"normalized": false,
|
| 17 |
+
"rstrip": false,
|
| 18 |
+
"single_word": false,
|
| 19 |
+
"special": true
|
| 20 |
+
},
|
| 21 |
+
"151645": {
|
| 22 |
+
"content": "<|im_end|>",
|
| 23 |
+
"lstrip": false,
|
| 24 |
+
"normalized": false,
|
| 25 |
+
"rstrip": false,
|
| 26 |
+
"single_word": false,
|
| 27 |
+
"special": true
|
| 28 |
+
},
|
| 29 |
+
"151646": {
|
| 30 |
+
"content": "<|object_ref_start|>",
|
| 31 |
+
"lstrip": false,
|
| 32 |
+
"normalized": false,
|
| 33 |
+
"rstrip": false,
|
| 34 |
+
"single_word": false,
|
| 35 |
+
"special": true
|
| 36 |
+
},
|
| 37 |
+
"151647": {
|
| 38 |
+
"content": "<|object_ref_end|>",
|
| 39 |
+
"lstrip": false,
|
| 40 |
+
"normalized": false,
|
| 41 |
+
"rstrip": false,
|
| 42 |
+
"single_word": false,
|
| 43 |
+
"special": true
|
| 44 |
+
},
|
| 45 |
+
"151648": {
|
| 46 |
+
"content": "<|box_start|>",
|
| 47 |
+
"lstrip": false,
|
| 48 |
+
"normalized": false,
|
| 49 |
+
"rstrip": false,
|
| 50 |
+
"single_word": false,
|
| 51 |
+
"special": true
|
| 52 |
+
},
|
| 53 |
+
"151649": {
|
| 54 |
+
"content": "<|box_end|>",
|
| 55 |
+
"lstrip": false,
|
| 56 |
+
"normalized": false,
|
| 57 |
+
"rstrip": false,
|
| 58 |
+
"single_word": false,
|
| 59 |
+
"special": true
|
| 60 |
+
},
|
| 61 |
+
"151650": {
|
| 62 |
+
"content": "<|quad_start|>",
|
| 63 |
+
"lstrip": false,
|
| 64 |
+
"normalized": false,
|
| 65 |
+
"rstrip": false,
|
| 66 |
+
"single_word": false,
|
| 67 |
+
"special": true
|
| 68 |
+
},
|
| 69 |
+
"151651": {
|
| 70 |
+
"content": "<|quad_end|>",
|
| 71 |
+
"lstrip": false,
|
| 72 |
+
"normalized": false,
|
| 73 |
+
"rstrip": false,
|
| 74 |
+
"single_word": false,
|
| 75 |
+
"special": true
|
| 76 |
+
},
|
| 77 |
+
"151652": {
|
| 78 |
+
"content": "<|vision_start|>",
|
| 79 |
+
"lstrip": false,
|
| 80 |
+
"normalized": false,
|
| 81 |
+
"rstrip": false,
|
| 82 |
+
"single_word": false,
|
| 83 |
+
"special": true
|
| 84 |
+
},
|
| 85 |
+
"151653": {
|
| 86 |
+
"content": "<|vision_end|>",
|
| 87 |
+
"lstrip": false,
|
| 88 |
+
"normalized": false,
|
| 89 |
+
"rstrip": false,
|
| 90 |
+
"single_word": false,
|
| 91 |
+
"special": true
|
| 92 |
+
},
|
| 93 |
+
"151654": {
|
| 94 |
+
"content": "<|vision_pad|>",
|
| 95 |
+
"lstrip": false,
|
| 96 |
+
"normalized": false,
|
| 97 |
+
"rstrip": false,
|
| 98 |
+
"single_word": false,
|
| 99 |
+
"special": true
|
| 100 |
+
},
|
| 101 |
+
"151655": {
|
| 102 |
+
"content": "<|image_pad|>",
|
| 103 |
+
"lstrip": false,
|
| 104 |
+
"normalized": false,
|
| 105 |
+
"rstrip": false,
|
| 106 |
+
"single_word": false,
|
| 107 |
+
"special": true
|
| 108 |
+
},
|
| 109 |
+
"151656": {
|
| 110 |
+
"content": "<|video_pad|>",
|
| 111 |
+
"lstrip": false,
|
| 112 |
+
"normalized": false,
|
| 113 |
+
"rstrip": false,
|
| 114 |
+
"single_word": false,
|
| 115 |
+
"special": true
|
| 116 |
+
},
|
| 117 |
+
"151657": {
|
| 118 |
+
"content": "<tool_call>",
|
| 119 |
+
"lstrip": false,
|
| 120 |
+
"normalized": false,
|
| 121 |
+
"rstrip": false,
|
| 122 |
+
"single_word": false,
|
| 123 |
+
"special": false
|
| 124 |
+
},
|
| 125 |
+
"151658": {
|
| 126 |
+
"content": "</tool_call>",
|
| 127 |
+
"lstrip": false,
|
| 128 |
+
"normalized": false,
|
| 129 |
+
"rstrip": false,
|
| 130 |
+
"single_word": false,
|
| 131 |
+
"special": false
|
| 132 |
+
},
|
| 133 |
+
"151659": {
|
| 134 |
+
"content": "<|fim_prefix|>",
|
| 135 |
+
"lstrip": false,
|
| 136 |
+
"normalized": false,
|
| 137 |
+
"rstrip": false,
|
| 138 |
+
"single_word": false,
|
| 139 |
+
"special": false
|
| 140 |
+
},
|
| 141 |
+
"151660": {
|
| 142 |
+
"content": "<|fim_middle|>",
|
| 143 |
+
"lstrip": false,
|
| 144 |
+
"normalized": false,
|
| 145 |
+
"rstrip": false,
|
| 146 |
+
"single_word": false,
|
| 147 |
+
"special": false
|
| 148 |
+
},
|
| 149 |
+
"151661": {
|
| 150 |
+
"content": "<|fim_suffix|>",
|
| 151 |
+
"lstrip": false,
|
| 152 |
+
"normalized": false,
|
| 153 |
+
"rstrip": false,
|
| 154 |
+
"single_word": false,
|
| 155 |
+
"special": false
|
| 156 |
+
},
|
| 157 |
+
"151662": {
|
| 158 |
+
"content": "<|fim_pad|>",
|
| 159 |
+
"lstrip": false,
|
| 160 |
+
"normalized": false,
|
| 161 |
+
"rstrip": false,
|
| 162 |
+
"single_word": false,
|
| 163 |
+
"special": false
|
| 164 |
+
},
|
| 165 |
+
"151663": {
|
| 166 |
+
"content": "<|repo_name|>",
|
| 167 |
+
"lstrip": false,
|
| 168 |
+
"normalized": false,
|
| 169 |
+
"rstrip": false,
|
| 170 |
+
"single_word": false,
|
| 171 |
+
"special": false
|
| 172 |
+
},
|
| 173 |
+
"151664": {
|
| 174 |
+
"content": "<|file_sep|>",
|
| 175 |
+
"lstrip": false,
|
| 176 |
+
"normalized": false,
|
| 177 |
+
"rstrip": false,
|
| 178 |
+
"single_word": false,
|
| 179 |
+
"special": false
|
| 180 |
+
},
|
| 181 |
+
"151665": {
|
| 182 |
+
"content": "<tool_response>",
|
| 183 |
+
"lstrip": false,
|
| 184 |
+
"normalized": false,
|
| 185 |
+
"rstrip": false,
|
| 186 |
+
"single_word": false,
|
| 187 |
+
"special": false
|
| 188 |
+
},
|
| 189 |
+
"151666": {
|
| 190 |
+
"content": "</tool_response>",
|
| 191 |
+
"lstrip": false,
|
| 192 |
+
"normalized": false,
|
| 193 |
+
"rstrip": false,
|
| 194 |
+
"single_word": false,
|
| 195 |
+
"special": false
|
| 196 |
+
},
|
| 197 |
+
"151667": {
|
| 198 |
+
"content": "<think>",
|
| 199 |
+
"lstrip": false,
|
| 200 |
+
"normalized": false,
|
| 201 |
+
"rstrip": false,
|
| 202 |
+
"single_word": false,
|
| 203 |
+
"special": false
|
| 204 |
+
},
|
| 205 |
+
"151668": {
|
| 206 |
+
"content": "</think>",
|
| 207 |
+
"lstrip": false,
|
| 208 |
+
"normalized": false,
|
| 209 |
+
"rstrip": false,
|
| 210 |
+
"single_word": false,
|
| 211 |
+
"special": false
|
| 212 |
+
}
|
| 213 |
+
},
|
| 214 |
+
"additional_special_tokens": [
|
| 215 |
+
"<|im_start|>",
|
| 216 |
+
"<|im_end|>",
|
| 217 |
+
"<|object_ref_start|>",
|
| 218 |
+
"<|object_ref_end|>",
|
| 219 |
+
"<|box_start|>",
|
| 220 |
+
"<|box_end|>",
|
| 221 |
+
"<|quad_start|>",
|
| 222 |
+
"<|quad_end|>",
|
| 223 |
+
"<|vision_start|>",
|
| 224 |
+
"<|vision_end|>",
|
| 225 |
+
"<|vision_pad|>",
|
| 226 |
+
"<|image_pad|>",
|
| 227 |
+
"<|video_pad|>",
|
| 228 |
+
"<d>",
|
| 229 |
+
"</d>",
|
| 230 |
+
"<|cutoff|>",
|
| 231 |
+
"<|lyrics_start|>",
|
| 232 |
+
"<|lyrics_end|>",
|
| 233 |
+
"<|caption_start|>",
|
| 234 |
+
"<|caption_end|>"
|
| 235 |
+
],
|
| 236 |
+
"bos_token": null,
|
| 237 |
+
"chat_template": "{%- if tools %}\n {{- '<|im_start|>system\\n' }}\n {%- if messages[0].role == 'system' %}\n {%- if messages[0].content is string %}\n {{- messages[0].content }}\n {%- else %}\n {%- for content in messages[0].content %}\n {%- if 'text' in content %}\n {{- content.text }}\n {%- endif %}\n {%- endfor %}\n {%- endif %}\n {{- '\\n\\n' }}\n {%- endif %}\n {{- \"# Tools\\n\\nYou may call one or more functions to assist with the user query.\\n\\nYou are provided with function signatures within <tools></tools> XML tags:\\n<tools>\" }}\n {%- for tool in tools %}\n {{- \"\\n\" }}\n {{- tool | tojson }}\n {%- endfor %}\n {{- \"\\n</tools>\\n\\nFor each function call, return a json object with function name and arguments within <tool_call></tool_call> XML tags:\\n<tool_call>\\n{\\\"name\\\": <function-name>, \\\"arguments\\\": <args-json-object>}\\n</tool_call><|im_end|>\\n\" }}\n{%- else %}\n {%- if messages[0].role == 'system' %}\n {{- '<|im_start|>system\\n' }}\n {%- if messages[0].content is string %}\n {{- messages[0].content }}\n {%- else %}\n {%- for content in messages[0].content %}\n {%- if 'text' in content %}\n {{- content.text }}\n {%- endif %}\n {%- endfor %}\n {%- endif %}\n {{- '<|im_end|>\\n' }}\n {%- endif %}\n{%- endif %}\n{%- set image_count = namespace(value=0) %}\n{%- set video_count = namespace(value=0) %}\n{%- for message in messages %}\n {%- if message.role == \"user\" %}\n {{- '<|im_start|>' + message.role + '\\n' }}\n {%- if message.content is string %}\n {{- message.content }}\n {%- else %}\n {%- for content in message.content %}\n {%- if content.type == 'image' or 'image' in content or 'image_url' in content %}\n {%- set image_count.value = image_count.value + 1 %}\n {%- if add_vision_id %}Picture {{ image_count.value }}: {% endif -%}\n <|vision_start|><|image_pad|><|vision_end|>\n {%- elif content.type == 'video' or 'video' in content %}\n {%- set video_count.value = video_count.value + 1 %}\n {%- if add_vision_id %}Video {{ video_count.value }}: {% endif -%}\n <|vision_start|><|video_pad|><|vision_end|>\n {%- elif 'text' in content %}\n {{- content.text }}\n {%- endif %}\n {%- endfor %}\n {%- endif %}\n {{- '<|im_end|>\\n' }}\n {%- elif message.role == \"assistant\" %}\n {{- '<|im_start|>' + message.role + '\\n' }}\n {%- if message.content is string %}\n {{- message.content }}\n {%- else %}\n {%- for content_item in message.content %}\n {%- if 'text' in content_item %}\n {{- content_item.text }}\n {%- endif %}\n {%- endfor %}\n {%- endif %}\n {%- if message.tool_calls %}\n {%- for tool_call in message.tool_calls %}\n {%- if (loop.first and message.content) or (not loop.first) %}\n {{- '\\n' }}\n {%- endif %}\n {%- if tool_call.function %}\n {%- set tool_call = tool_call.function %}\n {%- endif %}\n {{- '<tool_call>\\n{\"name\": \"' }}\n {{- tool_call.name }}\n {{- '\", \"arguments\": ' }}\n {%- if tool_call.arguments is string %}\n {{- tool_call.arguments }}\n {%- else %}\n {{- tool_call.arguments | tojson }}\n {%- endif %}\n {{- '}\\n</tool_call>' }}\n {%- endfor %}\n {%- endif %}\n {{- '<|im_end|>\\n' }}\n {%- elif message.role == \"tool\" %}\n {%- if loop.first or (messages[loop.index0 - 1].role != \"tool\") %}\n {{- '<|im_start|>user' }}\n {%- endif %}\n {{- '\\n<tool_response>\\n' }}\n {%- if message.content is string %}\n {{- message.content }}\n {%- else %}\n {%- for content in message.content %}\n {%- if content.type == 'image' or 'image' in content or 'image_url' in content %}\n {%- set image_count.value = image_count.value + 1 %}\n {%- if add_vision_id %}Picture {{ image_count.value }}: {% endif -%}\n <|vision_start|><|image_pad|><|vision_end|>\n {%- elif content.type == 'video' or 'video' in content %}\n {%- set video_count.value = video_count.value + 1 %}\n {%- if add_vision_id %}Video {{ video_count.value }}: {% endif -%}\n <|vision_start|><|video_pad|><|vision_end|>\n {%- elif 'text' in content %}\n {{- content.text }}\n {%- endif %}\n {%- endfor %}\n {%- endif %}\n {{- '\\n</tool_response>' }}\n {%- if loop.last or (messages[loop.index0 + 1].role != \"tool\") %}\n {{- '<|im_end|>\\n' }}\n {%- endif %}\n {%- endif %}\n{%- endfor %}\n{%- if add_generation_prompt %}\n {{- '<|im_start|>assistant\\n' }}\n{%- endif %}\n",
|
| 238 |
+
"clean_up_tokenization_spaces": false,
|
| 239 |
+
"eos_token": "<|im_end|>",
|
| 240 |
+
"errors": "replace",
|
| 241 |
+
"model_max_length": 262144,
|
| 242 |
+
"pad_token": "<|endoftext|>",
|
| 243 |
+
"split_special_tokens": false,
|
| 244 |
+
"tokenizer_class": "Qwen2Tokenizer",
|
| 245 |
+
"unk_token": null
|
| 246 |
+
}
|
FL2VA/text_encoder/video_preprocessor_config.json
ADDED
|
@@ -0,0 +1,21 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"size": {
|
| 3 |
+
"longest_edge": 25165824,
|
| 4 |
+
"shortest_edge": 4096
|
| 5 |
+
},
|
| 6 |
+
"patch_size": 16,
|
| 7 |
+
"temporal_patch_size": 2,
|
| 8 |
+
"merge_size": 2,
|
| 9 |
+
"image_mean": [
|
| 10 |
+
0.5,
|
| 11 |
+
0.5,
|
| 12 |
+
0.5
|
| 13 |
+
],
|
| 14 |
+
"image_std": [
|
| 15 |
+
0.5,
|
| 16 |
+
0.5,
|
| 17 |
+
0.5
|
| 18 |
+
],
|
| 19 |
+
"processor_class": "Qwen3VLProcessor",
|
| 20 |
+
"video_processor_type": "Qwen3VLVideoProcessor"
|
| 21 |
+
}
|
FL2VA/text_encoder/vocab.json
ADDED
|
The diff for this file is too large to render.
See raw diff
|
|
|
FL2VA/tokenizer/merges.txt
ADDED
|
The diff for this file is too large to render.
See raw diff
|
|
|
FL2VA/tokenizer/tokenizer.json
ADDED
|
The diff for this file is too large to render.
See raw diff
|
|
|
FL2VA/tokenizer/tokenizer_config.json
ADDED
|
@@ -0,0 +1,246 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"add_bos_token": false,
|
| 3 |
+
"add_prefix_space": false,
|
| 4 |
+
"added_tokens_decoder": {
|
| 5 |
+
"151643": {
|
| 6 |
+
"content": "<|endoftext|>",
|
| 7 |
+
"lstrip": false,
|
| 8 |
+
"normalized": false,
|
| 9 |
+
"rstrip": false,
|
| 10 |
+
"single_word": false,
|
| 11 |
+
"special": true
|
| 12 |
+
},
|
| 13 |
+
"151644": {
|
| 14 |
+
"content": "<|im_start|>",
|
| 15 |
+
"lstrip": false,
|
| 16 |
+
"normalized": false,
|
| 17 |
+
"rstrip": false,
|
| 18 |
+
"single_word": false,
|
| 19 |
+
"special": true
|
| 20 |
+
},
|
| 21 |
+
"151645": {
|
| 22 |
+
"content": "<|im_end|>",
|
| 23 |
+
"lstrip": false,
|
| 24 |
+
"normalized": false,
|
| 25 |
+
"rstrip": false,
|
| 26 |
+
"single_word": false,
|
| 27 |
+
"special": true
|
| 28 |
+
},
|
| 29 |
+
"151646": {
|
| 30 |
+
"content": "<|object_ref_start|>",
|
| 31 |
+
"lstrip": false,
|
| 32 |
+
"normalized": false,
|
| 33 |
+
"rstrip": false,
|
| 34 |
+
"single_word": false,
|
| 35 |
+
"special": true
|
| 36 |
+
},
|
| 37 |
+
"151647": {
|
| 38 |
+
"content": "<|object_ref_end|>",
|
| 39 |
+
"lstrip": false,
|
| 40 |
+
"normalized": false,
|
| 41 |
+
"rstrip": false,
|
| 42 |
+
"single_word": false,
|
| 43 |
+
"special": true
|
| 44 |
+
},
|
| 45 |
+
"151648": {
|
| 46 |
+
"content": "<|box_start|>",
|
| 47 |
+
"lstrip": false,
|
| 48 |
+
"normalized": false,
|
| 49 |
+
"rstrip": false,
|
| 50 |
+
"single_word": false,
|
| 51 |
+
"special": true
|
| 52 |
+
},
|
| 53 |
+
"151649": {
|
| 54 |
+
"content": "<|box_end|>",
|
| 55 |
+
"lstrip": false,
|
| 56 |
+
"normalized": false,
|
| 57 |
+
"rstrip": false,
|
| 58 |
+
"single_word": false,
|
| 59 |
+
"special": true
|
| 60 |
+
},
|
| 61 |
+
"151650": {
|
| 62 |
+
"content": "<|quad_start|>",
|
| 63 |
+
"lstrip": false,
|
| 64 |
+
"normalized": false,
|
| 65 |
+
"rstrip": false,
|
| 66 |
+
"single_word": false,
|
| 67 |
+
"special": true
|
| 68 |
+
},
|
| 69 |
+
"151651": {
|
| 70 |
+
"content": "<|quad_end|>",
|
| 71 |
+
"lstrip": false,
|
| 72 |
+
"normalized": false,
|
| 73 |
+
"rstrip": false,
|
| 74 |
+
"single_word": false,
|
| 75 |
+
"special": true
|
| 76 |
+
},
|
| 77 |
+
"151652": {
|
| 78 |
+
"content": "<|vision_start|>",
|
| 79 |
+
"lstrip": false,
|
| 80 |
+
"normalized": false,
|
| 81 |
+
"rstrip": false,
|
| 82 |
+
"single_word": false,
|
| 83 |
+
"special": true
|
| 84 |
+
},
|
| 85 |
+
"151653": {
|
| 86 |
+
"content": "<|vision_end|>",
|
| 87 |
+
"lstrip": false,
|
| 88 |
+
"normalized": false,
|
| 89 |
+
"rstrip": false,
|
| 90 |
+
"single_word": false,
|
| 91 |
+
"special": true
|
| 92 |
+
},
|
| 93 |
+
"151654": {
|
| 94 |
+
"content": "<|vision_pad|>",
|
| 95 |
+
"lstrip": false,
|
| 96 |
+
"normalized": false,
|
| 97 |
+
"rstrip": false,
|
| 98 |
+
"single_word": false,
|
| 99 |
+
"special": true
|
| 100 |
+
},
|
| 101 |
+
"151655": {
|
| 102 |
+
"content": "<|image_pad|>",
|
| 103 |
+
"lstrip": false,
|
| 104 |
+
"normalized": false,
|
| 105 |
+
"rstrip": false,
|
| 106 |
+
"single_word": false,
|
| 107 |
+
"special": true
|
| 108 |
+
},
|
| 109 |
+
"151656": {
|
| 110 |
+
"content": "<|video_pad|>",
|
| 111 |
+
"lstrip": false,
|
| 112 |
+
"normalized": false,
|
| 113 |
+
"rstrip": false,
|
| 114 |
+
"single_word": false,
|
| 115 |
+
"special": true
|
| 116 |
+
},
|
| 117 |
+
"151657": {
|
| 118 |
+
"content": "<tool_call>",
|
| 119 |
+
"lstrip": false,
|
| 120 |
+
"normalized": false,
|
| 121 |
+
"rstrip": false,
|
| 122 |
+
"single_word": false,
|
| 123 |
+
"special": false
|
| 124 |
+
},
|
| 125 |
+
"151658": {
|
| 126 |
+
"content": "</tool_call>",
|
| 127 |
+
"lstrip": false,
|
| 128 |
+
"normalized": false,
|
| 129 |
+
"rstrip": false,
|
| 130 |
+
"single_word": false,
|
| 131 |
+
"special": false
|
| 132 |
+
},
|
| 133 |
+
"151659": {
|
| 134 |
+
"content": "<|fim_prefix|>",
|
| 135 |
+
"lstrip": false,
|
| 136 |
+
"normalized": false,
|
| 137 |
+
"rstrip": false,
|
| 138 |
+
"single_word": false,
|
| 139 |
+
"special": false
|
| 140 |
+
},
|
| 141 |
+
"151660": {
|
| 142 |
+
"content": "<|fim_middle|>",
|
| 143 |
+
"lstrip": false,
|
| 144 |
+
"normalized": false,
|
| 145 |
+
"rstrip": false,
|
| 146 |
+
"single_word": false,
|
| 147 |
+
"special": false
|
| 148 |
+
},
|
| 149 |
+
"151661": {
|
| 150 |
+
"content": "<|fim_suffix|>",
|
| 151 |
+
"lstrip": false,
|
| 152 |
+
"normalized": false,
|
| 153 |
+
"rstrip": false,
|
| 154 |
+
"single_word": false,
|
| 155 |
+
"special": false
|
| 156 |
+
},
|
| 157 |
+
"151662": {
|
| 158 |
+
"content": "<|fim_pad|>",
|
| 159 |
+
"lstrip": false,
|
| 160 |
+
"normalized": false,
|
| 161 |
+
"rstrip": false,
|
| 162 |
+
"single_word": false,
|
| 163 |
+
"special": false
|
| 164 |
+
},
|
| 165 |
+
"151663": {
|
| 166 |
+
"content": "<|repo_name|>",
|
| 167 |
+
"lstrip": false,
|
| 168 |
+
"normalized": false,
|
| 169 |
+
"rstrip": false,
|
| 170 |
+
"single_word": false,
|
| 171 |
+
"special": false
|
| 172 |
+
},
|
| 173 |
+
"151664": {
|
| 174 |
+
"content": "<|file_sep|>",
|
| 175 |
+
"lstrip": false,
|
| 176 |
+
"normalized": false,
|
| 177 |
+
"rstrip": false,
|
| 178 |
+
"single_word": false,
|
| 179 |
+
"special": false
|
| 180 |
+
},
|
| 181 |
+
"151665": {
|
| 182 |
+
"content": "<tool_response>",
|
| 183 |
+
"lstrip": false,
|
| 184 |
+
"normalized": false,
|
| 185 |
+
"rstrip": false,
|
| 186 |
+
"single_word": false,
|
| 187 |
+
"special": false
|
| 188 |
+
},
|
| 189 |
+
"151666": {
|
| 190 |
+
"content": "</tool_response>",
|
| 191 |
+
"lstrip": false,
|
| 192 |
+
"normalized": false,
|
| 193 |
+
"rstrip": false,
|
| 194 |
+
"single_word": false,
|
| 195 |
+
"special": false
|
| 196 |
+
},
|
| 197 |
+
"151667": {
|
| 198 |
+
"content": "<think>",
|
| 199 |
+
"lstrip": false,
|
| 200 |
+
"normalized": false,
|
| 201 |
+
"rstrip": false,
|
| 202 |
+
"single_word": false,
|
| 203 |
+
"special": false
|
| 204 |
+
},
|
| 205 |
+
"151668": {
|
| 206 |
+
"content": "</think>",
|
| 207 |
+
"lstrip": false,
|
| 208 |
+
"normalized": false,
|
| 209 |
+
"rstrip": false,
|
| 210 |
+
"single_word": false,
|
| 211 |
+
"special": false
|
| 212 |
+
}
|
| 213 |
+
},
|
| 214 |
+
"additional_special_tokens": [
|
| 215 |
+
"<|im_start|>",
|
| 216 |
+
"<|im_end|>",
|
| 217 |
+
"<|object_ref_start|>",
|
| 218 |
+
"<|object_ref_end|>",
|
| 219 |
+
"<|box_start|>",
|
| 220 |
+
"<|box_end|>",
|
| 221 |
+
"<|quad_start|>",
|
| 222 |
+
"<|quad_end|>",
|
| 223 |
+
"<|vision_start|>",
|
| 224 |
+
"<|vision_end|>",
|
| 225 |
+
"<|vision_pad|>",
|
| 226 |
+
"<|image_pad|>",
|
| 227 |
+
"<|video_pad|>",
|
| 228 |
+
"<d>",
|
| 229 |
+
"</d>",
|
| 230 |
+
"<|cutoff|>",
|
| 231 |
+
"<|lyrics_start|>",
|
| 232 |
+
"<|lyrics_end|>",
|
| 233 |
+
"<|caption_start|>",
|
| 234 |
+
"<|caption_end|>"
|
| 235 |
+
],
|
| 236 |
+
"bos_token": null,
|
| 237 |
+
"chat_template": "{%- if tools %}\n {{- '<|im_start|>system\\n' }}\n {%- if messages[0].role == 'system' %}\n {%- if messages[0].content is string %}\n {{- messages[0].content }}\n {%- else %}\n {%- for content in messages[0].content %}\n {%- if 'text' in content %}\n {{- content.text }}\n {%- endif %}\n {%- endfor %}\n {%- endif %}\n {{- '\\n\\n' }}\n {%- endif %}\n {{- \"# Tools\\n\\nYou may call one or more functions to assist with the user query.\\n\\nYou are provided with function signatures within <tools></tools> XML tags:\\n<tools>\" }}\n {%- for tool in tools %}\n {{- \"\\n\" }}\n {{- tool | tojson }}\n {%- endfor %}\n {{- \"\\n</tools>\\n\\nFor each function call, return a json object with function name and arguments within <tool_call></tool_call> XML tags:\\n<tool_call>\\n{\\\"name\\\": <function-name>, \\\"arguments\\\": <args-json-object>}\\n</tool_call><|im_end|>\\n\" }}\n{%- else %}\n {%- if messages[0].role == 'system' %}\n {{- '<|im_start|>system\\n' }}\n {%- if messages[0].content is string %}\n {{- messages[0].content }}\n {%- else %}\n {%- for content in messages[0].content %}\n {%- if 'text' in content %}\n {{- content.text }}\n {%- endif %}\n {%- endfor %}\n {%- endif %}\n {{- '<|im_end|>\\n' }}\n {%- endif %}\n{%- endif %}\n{%- set image_count = namespace(value=0) %}\n{%- set video_count = namespace(value=0) %}\n{%- for message in messages %}\n {%- if message.role == \"user\" %}\n {{- '<|im_start|>' + message.role + '\\n' }}\n {%- if message.content is string %}\n {{- message.content }}\n {%- else %}\n {%- for content in message.content %}\n {%- if content.type == 'image' or 'image' in content or 'image_url' in content %}\n {%- set image_count.value = image_count.value + 1 %}\n {%- if add_vision_id %}Picture {{ image_count.value }}: {% endif -%}\n <|vision_start|><|image_pad|><|vision_end|>\n {%- elif content.type == 'video' or 'video' in content %}\n {%- set video_count.value = video_count.value + 1 %}\n {%- if add_vision_id %}Video {{ video_count.value }}: {% endif -%}\n <|vision_start|><|video_pad|><|vision_end|>\n {%- elif 'text' in content %}\n {{- content.text }}\n {%- endif %}\n {%- endfor %}\n {%- endif %}\n {{- '<|im_end|>\\n' }}\n {%- elif message.role == \"assistant\" %}\n {{- '<|im_start|>' + message.role + '\\n' }}\n {%- if message.content is string %}\n {{- message.content }}\n {%- else %}\n {%- for content_item in message.content %}\n {%- if 'text' in content_item %}\n {{- content_item.text }}\n {%- endif %}\n {%- endfor %}\n {%- endif %}\n {%- if message.tool_calls %}\n {%- for tool_call in message.tool_calls %}\n {%- if (loop.first and message.content) or (not loop.first) %}\n {{- '\\n' }}\n {%- endif %}\n {%- if tool_call.function %}\n {%- set tool_call = tool_call.function %}\n {%- endif %}\n {{- '<tool_call>\\n{\"name\": \"' }}\n {{- tool_call.name }}\n {{- '\", \"arguments\": ' }}\n {%- if tool_call.arguments is string %}\n {{- tool_call.arguments }}\n {%- else %}\n {{- tool_call.arguments | tojson }}\n {%- endif %}\n {{- '}\\n</tool_call>' }}\n {%- endfor %}\n {%- endif %}\n {{- '<|im_end|>\\n' }}\n {%- elif message.role == \"tool\" %}\n {%- if loop.first or (messages[loop.index0 - 1].role != \"tool\") %}\n {{- '<|im_start|>user' }}\n {%- endif %}\n {{- '\\n<tool_response>\\n' }}\n {%- if message.content is string %}\n {{- message.content }}\n {%- else %}\n {%- for content in message.content %}\n {%- if content.type == 'image' or 'image' in content or 'image_url' in content %}\n {%- set image_count.value = image_count.value + 1 %}\n {%- if add_vision_id %}Picture {{ image_count.value }}: {% endif -%}\n <|vision_start|><|image_pad|><|vision_end|>\n {%- elif content.type == 'video' or 'video' in content %}\n {%- set video_count.value = video_count.value + 1 %}\n {%- if add_vision_id %}Video {{ video_count.value }}: {% endif -%}\n <|vision_start|><|video_pad|><|vision_end|>\n {%- elif 'text' in content %}\n {{- content.text }}\n {%- endif %}\n {%- endfor %}\n {%- endif %}\n {{- '\\n</tool_response>' }}\n {%- if loop.last or (messages[loop.index0 + 1].role != \"tool\") %}\n {{- '<|im_end|>\\n' }}\n {%- endif %}\n {%- endif %}\n{%- endfor %}\n{%- if add_generation_prompt %}\n {{- '<|im_start|>assistant\\n' }}\n{%- endif %}\n",
|
| 238 |
+
"clean_up_tokenization_spaces": false,
|
| 239 |
+
"eos_token": "<|im_end|>",
|
| 240 |
+
"errors": "replace",
|
| 241 |
+
"model_max_length": 262144,
|
| 242 |
+
"pad_token": "<|endoftext|>",
|
| 243 |
+
"split_special_tokens": false,
|
| 244 |
+
"tokenizer_class": "Qwen2Tokenizer",
|
| 245 |
+
"unk_token": null
|
| 246 |
+
}
|
FL2VA/tokenizer/vocab.json
ADDED
|
The diff for this file is too large to render.
See raw diff
|
|
|
FL2VA/transformer/config.json
ADDED
|
@@ -0,0 +1,27 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"_class_name": "MiniMaxH3DiTModel",
|
| 3 |
+
"_diffusers_version": "0.32.2",
|
| 4 |
+
"hidden_size": 5376,
|
| 5 |
+
"num_layers": 50,
|
| 6 |
+
"token_refiner_num_layers": 2,
|
| 7 |
+
"num_attention_heads": 56,
|
| 8 |
+
"attention_head_dim": 128,
|
| 9 |
+
"ffn_hidden_size": 14336,
|
| 10 |
+
"latents_dim": 24,
|
| 11 |
+
"audio_latents_dim": 32,
|
| 12 |
+
"patch_size": [
|
| 13 |
+
1,
|
| 14 |
+
2,
|
| 15 |
+
2
|
| 16 |
+
],
|
| 17 |
+
"text_dim": 5120,
|
| 18 |
+
"timestep_input_dim": 256,
|
| 19 |
+
"time_embed_hidden_size": 5376,
|
| 20 |
+
"time_embed_dim": 2688,
|
| 21 |
+
"adaln_out_features": 96768,
|
| 22 |
+
"final_adaln_out_features": 10752,
|
| 23 |
+
"rope_inv_freq_len": 16,
|
| 24 |
+
"norm_eps": 1e-05,
|
| 25 |
+
"qk_norm_eps": 1e-05,
|
| 26 |
+
"final_norm_eps": 1e-05
|
| 27 |
+
}
|
FL2VA/transformer/model-00013-of-00013.safetensors
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:8bfd852d5817e9836de1d3ec8dbac1c5446b167568b717371f08282b22291aa2
|
| 3 |
+
size 4242305176
|
FL2VA/transformer/model.safetensors.index.json
ADDED
|
@@ -0,0 +1,542 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"metadata": {
|
| 3 |
+
"total_size": 66280430144
|
| 4 |
+
},
|
| 5 |
+
"weight_map": {
|
| 6 |
+
"audio_patch_proj.bias": "model-00001-of-00013.safetensors",
|
| 7 |
+
"audio_patch_proj.weight": "model-00001-of-00013.safetensors",
|
| 8 |
+
"blocks.0.adaln_proj.linear.bias": "model-00001-of-00013.safetensors",
|
| 9 |
+
"blocks.0.adaln_proj.linear.weight": "model-00001-of-00013.safetensors",
|
| 10 |
+
"blocks.0.attn.k_norm.weight": "model-00001-of-00013.safetensors",
|
| 11 |
+
"blocks.0.attn.out_proj.weight": "model-00001-of-00013.safetensors",
|
| 12 |
+
"blocks.0.attn.q_norm.weight": "model-00001-of-00013.safetensors",
|
| 13 |
+
"blocks.0.attn.qkv_proj.weight": "model-00001-of-00013.safetensors",
|
| 14 |
+
"blocks.0.mlp.fc1.weight": "model-00001-of-00013.safetensors",
|
| 15 |
+
"blocks.0.mlp.fc2.weight": "model-00001-of-00013.safetensors",
|
| 16 |
+
"blocks.0.norm1.weight": "model-00001-of-00013.safetensors",
|
| 17 |
+
"blocks.0.norm2.weight": "model-00001-of-00013.safetensors",
|
| 18 |
+
"blocks.1.adaln_proj.linear.bias": "model-00001-of-00013.safetensors",
|
| 19 |
+
"blocks.1.adaln_proj.linear.weight": "model-00001-of-00013.safetensors",
|
| 20 |
+
"blocks.1.attn.k_norm.weight": "model-00001-of-00013.safetensors",
|
| 21 |
+
"blocks.1.attn.out_proj.weight": "model-00001-of-00013.safetensors",
|
| 22 |
+
"blocks.1.attn.q_norm.weight": "model-00001-of-00013.safetensors",
|
| 23 |
+
"blocks.1.attn.qkv_proj.weight": "model-00001-of-00013.safetensors",
|
| 24 |
+
"blocks.1.mlp.fc1.weight": "model-00001-of-00013.safetensors",
|
| 25 |
+
"blocks.1.mlp.fc2.weight": "model-00001-of-00013.safetensors",
|
| 26 |
+
"blocks.1.norm1.weight": "model-00001-of-00013.safetensors",
|
| 27 |
+
"blocks.1.norm2.weight": "model-00001-of-00013.safetensors",
|
| 28 |
+
"blocks.10.adaln_proj.linear.bias": "model-00003-of-00013.safetensors",
|
| 29 |
+
"blocks.10.adaln_proj.linear.weight": "model-00003-of-00013.safetensors",
|
| 30 |
+
"blocks.10.attn.k_norm.weight": "model-00004-of-00013.safetensors",
|
| 31 |
+
"blocks.10.attn.out_proj.weight": "model-00004-of-00013.safetensors",
|
| 32 |
+
"blocks.10.attn.q_norm.weight": "model-00004-of-00013.safetensors",
|
| 33 |
+
"blocks.10.attn.qkv_proj.weight": "model-00004-of-00013.safetensors",
|
| 34 |
+
"blocks.10.mlp.fc1.weight": "model-00003-of-00013.safetensors",
|
| 35 |
+
"blocks.10.mlp.fc2.weight": "model-00003-of-00013.safetensors",
|
| 36 |
+
"blocks.10.norm1.weight": "model-00003-of-00013.safetensors",
|
| 37 |
+
"blocks.10.norm2.weight": "model-00003-of-00013.safetensors",
|
| 38 |
+
"blocks.11.adaln_proj.linear.bias": "model-00004-of-00013.safetensors",
|
| 39 |
+
"blocks.11.adaln_proj.linear.weight": "model-00004-of-00013.safetensors",
|
| 40 |
+
"blocks.11.attn.k_norm.weight": "model-00004-of-00013.safetensors",
|
| 41 |
+
"blocks.11.attn.out_proj.weight": "model-00004-of-00013.safetensors",
|
| 42 |
+
"blocks.11.attn.q_norm.weight": "model-00004-of-00013.safetensors",
|
| 43 |
+
"blocks.11.attn.qkv_proj.weight": "model-00004-of-00013.safetensors",
|
| 44 |
+
"blocks.11.mlp.fc1.weight": "model-00004-of-00013.safetensors",
|
| 45 |
+
"blocks.11.mlp.fc2.weight": "model-00004-of-00013.safetensors",
|
| 46 |
+
"blocks.11.norm1.weight": "model-00004-of-00013.safetensors",
|
| 47 |
+
"blocks.11.norm2.weight": "model-00004-of-00013.safetensors",
|
| 48 |
+
"blocks.12.adaln_proj.linear.bias": "model-00004-of-00013.safetensors",
|
| 49 |
+
"blocks.12.adaln_proj.linear.weight": "model-00004-of-00013.safetensors",
|
| 50 |
+
"blocks.12.attn.k_norm.weight": "model-00004-of-00013.safetensors",
|
| 51 |
+
"blocks.12.attn.out_proj.weight": "model-00004-of-00013.safetensors",
|
| 52 |
+
"blocks.12.attn.q_norm.weight": "model-00004-of-00013.safetensors",
|
| 53 |
+
"blocks.12.attn.qkv_proj.weight": "model-00004-of-00013.safetensors",
|
| 54 |
+
"blocks.12.mlp.fc1.weight": "model-00004-of-00013.safetensors",
|
| 55 |
+
"blocks.12.mlp.fc2.weight": "model-00004-of-00013.safetensors",
|
| 56 |
+
"blocks.12.norm1.weight": "model-00004-of-00013.safetensors",
|
| 57 |
+
"blocks.12.norm2.weight": "model-00004-of-00013.safetensors",
|
| 58 |
+
"blocks.13.adaln_proj.linear.bias": "model-00004-of-00013.safetensors",
|
| 59 |
+
"blocks.13.adaln_proj.linear.weight": "model-00004-of-00013.safetensors",
|
| 60 |
+
"blocks.13.attn.k_norm.weight": "model-00004-of-00013.safetensors",
|
| 61 |
+
"blocks.13.attn.out_proj.weight": "model-00004-of-00013.safetensors",
|
| 62 |
+
"blocks.13.attn.q_norm.weight": "model-00004-of-00013.safetensors",
|
| 63 |
+
"blocks.13.attn.qkv_proj.weight": "model-00004-of-00013.safetensors",
|
| 64 |
+
"blocks.13.mlp.fc1.weight": "model-00004-of-00013.safetensors",
|
| 65 |
+
"blocks.13.mlp.fc2.weight": "model-00004-of-00013.safetensors",
|
| 66 |
+
"blocks.13.norm1.weight": "model-00004-of-00013.safetensors",
|
| 67 |
+
"blocks.13.norm2.weight": "model-00004-of-00013.safetensors",
|
| 68 |
+
"blocks.14.adaln_proj.linear.bias": "model-00004-of-00013.safetensors",
|
| 69 |
+
"blocks.14.adaln_proj.linear.weight": "model-00004-of-00013.safetensors",
|
| 70 |
+
"blocks.14.attn.k_norm.weight": "model-00005-of-00013.safetensors",
|
| 71 |
+
"blocks.14.attn.out_proj.weight": "model-00005-of-00013.safetensors",
|
| 72 |
+
"blocks.14.attn.q_norm.weight": "model-00005-of-00013.safetensors",
|
| 73 |
+
"blocks.14.attn.qkv_proj.weight": "model-00005-of-00013.safetensors",
|
| 74 |
+
"blocks.14.mlp.fc1.weight": "model-00004-of-00013.safetensors",
|
| 75 |
+
"blocks.14.mlp.fc2.weight": "model-00004-of-00013.safetensors",
|
| 76 |
+
"blocks.14.norm1.weight": "model-00004-of-00013.safetensors",
|
| 77 |
+
"blocks.14.norm2.weight": "model-00004-of-00013.safetensors",
|
| 78 |
+
"blocks.15.adaln_proj.linear.bias": "model-00005-of-00013.safetensors",
|
| 79 |
+
"blocks.15.adaln_proj.linear.weight": "model-00005-of-00013.safetensors",
|
| 80 |
+
"blocks.15.attn.k_norm.weight": "model-00005-of-00013.safetensors",
|
| 81 |
+
"blocks.15.attn.out_proj.weight": "model-00005-of-00013.safetensors",
|
| 82 |
+
"blocks.15.attn.q_norm.weight": "model-00005-of-00013.safetensors",
|
| 83 |
+
"blocks.15.attn.qkv_proj.weight": "model-00005-of-00013.safetensors",
|
| 84 |
+
"blocks.15.mlp.fc1.weight": "model-00005-of-00013.safetensors",
|
| 85 |
+
"blocks.15.mlp.fc2.weight": "model-00005-of-00013.safetensors",
|
| 86 |
+
"blocks.15.norm1.weight": "model-00005-of-00013.safetensors",
|
| 87 |
+
"blocks.15.norm2.weight": "model-00005-of-00013.safetensors",
|
| 88 |
+
"blocks.16.adaln_proj.linear.bias": "model-00005-of-00013.safetensors",
|
| 89 |
+
"blocks.16.adaln_proj.linear.weight": "model-00005-of-00013.safetensors",
|
| 90 |
+
"blocks.16.attn.k_norm.weight": "model-00005-of-00013.safetensors",
|
| 91 |
+
"blocks.16.attn.out_proj.weight": "model-00005-of-00013.safetensors",
|
| 92 |
+
"blocks.16.attn.q_norm.weight": "model-00005-of-00013.safetensors",
|
| 93 |
+
"blocks.16.attn.qkv_proj.weight": "model-00005-of-00013.safetensors",
|
| 94 |
+
"blocks.16.mlp.fc1.weight": "model-00005-of-00013.safetensors",
|
| 95 |
+
"blocks.16.mlp.fc2.weight": "model-00005-of-00013.safetensors",
|
| 96 |
+
"blocks.16.norm1.weight": "model-00005-of-00013.safetensors",
|
| 97 |
+
"blocks.16.norm2.weight": "model-00005-of-00013.safetensors",
|
| 98 |
+
"blocks.17.adaln_proj.linear.bias": "model-00005-of-00013.safetensors",
|
| 99 |
+
"blocks.17.adaln_proj.linear.weight": "model-00005-of-00013.safetensors",
|
| 100 |
+
"blocks.17.attn.k_norm.weight": "model-00005-of-00013.safetensors",
|
| 101 |
+
"blocks.17.attn.out_proj.weight": "model-00005-of-00013.safetensors",
|
| 102 |
+
"blocks.17.attn.q_norm.weight": "model-00005-of-00013.safetensors",
|
| 103 |
+
"blocks.17.attn.qkv_proj.weight": "model-00005-of-00013.safetensors",
|
| 104 |
+
"blocks.17.mlp.fc1.weight": "model-00005-of-00013.safetensors",
|
| 105 |
+
"blocks.17.mlp.fc2.weight": "model-00005-of-00013.safetensors",
|
| 106 |
+
"blocks.17.norm1.weight": "model-00005-of-00013.safetensors",
|
| 107 |
+
"blocks.17.norm2.weight": "model-00005-of-00013.safetensors",
|
| 108 |
+
"blocks.18.adaln_proj.linear.bias": "model-00005-of-00013.safetensors",
|
| 109 |
+
"blocks.18.adaln_proj.linear.weight": "model-00005-of-00013.safetensors",
|
| 110 |
+
"blocks.18.attn.k_norm.weight": "model-00006-of-00013.safetensors",
|
| 111 |
+
"blocks.18.attn.out_proj.weight": "model-00006-of-00013.safetensors",
|
| 112 |
+
"blocks.18.attn.q_norm.weight": "model-00006-of-00013.safetensors",
|
| 113 |
+
"blocks.18.attn.qkv_proj.weight": "model-00006-of-00013.safetensors",
|
| 114 |
+
"blocks.18.mlp.fc1.weight": "model-00005-of-00013.safetensors",
|
| 115 |
+
"blocks.18.mlp.fc2.weight": "model-00005-of-00013.safetensors",
|
| 116 |
+
"blocks.18.norm1.weight": "model-00005-of-00013.safetensors",
|
| 117 |
+
"blocks.18.norm2.weight": "model-00005-of-00013.safetensors",
|
| 118 |
+
"blocks.19.adaln_proj.linear.bias": "model-00006-of-00013.safetensors",
|
| 119 |
+
"blocks.19.adaln_proj.linear.weight": "model-00006-of-00013.safetensors",
|
| 120 |
+
"blocks.19.attn.k_norm.weight": "model-00006-of-00013.safetensors",
|
| 121 |
+
"blocks.19.attn.out_proj.weight": "model-00006-of-00013.safetensors",
|
| 122 |
+
"blocks.19.attn.q_norm.weight": "model-00006-of-00013.safetensors",
|
| 123 |
+
"blocks.19.attn.qkv_proj.weight": "model-00006-of-00013.safetensors",
|
| 124 |
+
"blocks.19.mlp.fc1.weight": "model-00006-of-00013.safetensors",
|
| 125 |
+
"blocks.19.mlp.fc2.weight": "model-00006-of-00013.safetensors",
|
| 126 |
+
"blocks.19.norm1.weight": "model-00006-of-00013.safetensors",
|
| 127 |
+
"blocks.19.norm2.weight": "model-00006-of-00013.safetensors",
|
| 128 |
+
"blocks.2.adaln_proj.linear.bias": "model-00001-of-00013.safetensors",
|
| 129 |
+
"blocks.2.adaln_proj.linear.weight": "model-00001-of-00013.safetensors",
|
| 130 |
+
"blocks.2.attn.k_norm.weight": "model-00002-of-00013.safetensors",
|
| 131 |
+
"blocks.2.attn.out_proj.weight": "model-00002-of-00013.safetensors",
|
| 132 |
+
"blocks.2.attn.q_norm.weight": "model-00002-of-00013.safetensors",
|
| 133 |
+
"blocks.2.attn.qkv_proj.weight": "model-00002-of-00013.safetensors",
|
| 134 |
+
"blocks.2.mlp.fc1.weight": "model-00001-of-00013.safetensors",
|
| 135 |
+
"blocks.2.mlp.fc2.weight": "model-00001-of-00013.safetensors",
|
| 136 |
+
"blocks.2.norm1.weight": "model-00001-of-00013.safetensors",
|
| 137 |
+
"blocks.2.norm2.weight": "model-00001-of-00013.safetensors",
|
| 138 |
+
"blocks.20.adaln_proj.linear.bias": "model-00006-of-00013.safetensors",
|
| 139 |
+
"blocks.20.adaln_proj.linear.weight": "model-00006-of-00013.safetensors",
|
| 140 |
+
"blocks.20.attn.k_norm.weight": "model-00006-of-00013.safetensors",
|
| 141 |
+
"blocks.20.attn.out_proj.weight": "model-00006-of-00013.safetensors",
|
| 142 |
+
"blocks.20.attn.q_norm.weight": "model-00006-of-00013.safetensors",
|
| 143 |
+
"blocks.20.attn.qkv_proj.weight": "model-00006-of-00013.safetensors",
|
| 144 |
+
"blocks.20.mlp.fc1.weight": "model-00006-of-00013.safetensors",
|
| 145 |
+
"blocks.20.mlp.fc2.weight": "model-00006-of-00013.safetensors",
|
| 146 |
+
"blocks.20.norm1.weight": "model-00006-of-00013.safetensors",
|
| 147 |
+
"blocks.20.norm2.weight": "model-00006-of-00013.safetensors",
|
| 148 |
+
"blocks.21.adaln_proj.linear.bias": "model-00006-of-00013.safetensors",
|
| 149 |
+
"blocks.21.adaln_proj.linear.weight": "model-00006-of-00013.safetensors",
|
| 150 |
+
"blocks.21.attn.k_norm.weight": "model-00006-of-00013.safetensors",
|
| 151 |
+
"blocks.21.attn.out_proj.weight": "model-00006-of-00013.safetensors",
|
| 152 |
+
"blocks.21.attn.q_norm.weight": "model-00006-of-00013.safetensors",
|
| 153 |
+
"blocks.21.attn.qkv_proj.weight": "model-00006-of-00013.safetensors",
|
| 154 |
+
"blocks.21.mlp.fc1.weight": "model-00006-of-00013.safetensors",
|
| 155 |
+
"blocks.21.mlp.fc2.weight": "model-00006-of-00013.safetensors",
|
| 156 |
+
"blocks.21.norm1.weight": "model-00006-of-00013.safetensors",
|
| 157 |
+
"blocks.21.norm2.weight": "model-00006-of-00013.safetensors",
|
| 158 |
+
"blocks.22.adaln_proj.linear.bias": "model-00006-of-00013.safetensors",
|
| 159 |
+
"blocks.22.adaln_proj.linear.weight": "model-00006-of-00013.safetensors",
|
| 160 |
+
"blocks.22.attn.k_norm.weight": "model-00007-of-00013.safetensors",
|
| 161 |
+
"blocks.22.attn.out_proj.weight": "model-00007-of-00013.safetensors",
|
| 162 |
+
"blocks.22.attn.q_norm.weight": "model-00007-of-00013.safetensors",
|
| 163 |
+
"blocks.22.attn.qkv_proj.weight": "model-00007-of-00013.safetensors",
|
| 164 |
+
"blocks.22.mlp.fc1.weight": "model-00006-of-00013.safetensors",
|
| 165 |
+
"blocks.22.mlp.fc2.weight": "model-00006-of-00013.safetensors",
|
| 166 |
+
"blocks.22.norm1.weight": "model-00006-of-00013.safetensors",
|
| 167 |
+
"blocks.22.norm2.weight": "model-00006-of-00013.safetensors",
|
| 168 |
+
"blocks.23.adaln_proj.linear.bias": "model-00007-of-00013.safetensors",
|
| 169 |
+
"blocks.23.adaln_proj.linear.weight": "model-00007-of-00013.safetensors",
|
| 170 |
+
"blocks.23.attn.k_norm.weight": "model-00007-of-00013.safetensors",
|
| 171 |
+
"blocks.23.attn.out_proj.weight": "model-00007-of-00013.safetensors",
|
| 172 |
+
"blocks.23.attn.q_norm.weight": "model-00007-of-00013.safetensors",
|
| 173 |
+
"blocks.23.attn.qkv_proj.weight": "model-00007-of-00013.safetensors",
|
| 174 |
+
"blocks.23.mlp.fc1.weight": "model-00007-of-00013.safetensors",
|
| 175 |
+
"blocks.23.mlp.fc2.weight": "model-00007-of-00013.safetensors",
|
| 176 |
+
"blocks.23.norm1.weight": "model-00007-of-00013.safetensors",
|
| 177 |
+
"blocks.23.norm2.weight": "model-00007-of-00013.safetensors",
|
| 178 |
+
"blocks.24.adaln_proj.linear.bias": "model-00007-of-00013.safetensors",
|
| 179 |
+
"blocks.24.adaln_proj.linear.weight": "model-00007-of-00013.safetensors",
|
| 180 |
+
"blocks.24.attn.k_norm.weight": "model-00007-of-00013.safetensors",
|
| 181 |
+
"blocks.24.attn.out_proj.weight": "model-00007-of-00013.safetensors",
|
| 182 |
+
"blocks.24.attn.q_norm.weight": "model-00007-of-00013.safetensors",
|
| 183 |
+
"blocks.24.attn.qkv_proj.weight": "model-00007-of-00013.safetensors",
|
| 184 |
+
"blocks.24.mlp.fc1.weight": "model-00007-of-00013.safetensors",
|
| 185 |
+
"blocks.24.mlp.fc2.weight": "model-00007-of-00013.safetensors",
|
| 186 |
+
"blocks.24.norm1.weight": "model-00007-of-00013.safetensors",
|
| 187 |
+
"blocks.24.norm2.weight": "model-00007-of-00013.safetensors",
|
| 188 |
+
"blocks.25.adaln_proj.linear.bias": "model-00007-of-00013.safetensors",
|
| 189 |
+
"blocks.25.adaln_proj.linear.weight": "model-00007-of-00013.safetensors",
|
| 190 |
+
"blocks.25.attn.k_norm.weight": "model-00007-of-00013.safetensors",
|
| 191 |
+
"blocks.25.attn.out_proj.weight": "model-00007-of-00013.safetensors",
|
| 192 |
+
"blocks.25.attn.q_norm.weight": "model-00007-of-00013.safetensors",
|
| 193 |
+
"blocks.25.attn.qkv_proj.weight": "model-00007-of-00013.safetensors",
|
| 194 |
+
"blocks.25.mlp.fc1.weight": "model-00007-of-00013.safetensors",
|
| 195 |
+
"blocks.25.mlp.fc2.weight": "model-00007-of-00013.safetensors",
|
| 196 |
+
"blocks.25.norm1.weight": "model-00007-of-00013.safetensors",
|
| 197 |
+
"blocks.25.norm2.weight": "model-00007-of-00013.safetensors",
|
| 198 |
+
"blocks.26.adaln_proj.linear.bias": "model-00007-of-00013.safetensors",
|
| 199 |
+
"blocks.26.adaln_proj.linear.weight": "model-00007-of-00013.safetensors",
|
| 200 |
+
"blocks.26.attn.k_norm.weight": "model-00008-of-00013.safetensors",
|
| 201 |
+
"blocks.26.attn.out_proj.weight": "model-00008-of-00013.safetensors",
|
| 202 |
+
"blocks.26.attn.q_norm.weight": "model-00008-of-00013.safetensors",
|
| 203 |
+
"blocks.26.attn.qkv_proj.weight": "model-00008-of-00013.safetensors",
|
| 204 |
+
"blocks.26.mlp.fc1.weight": "model-00007-of-00013.safetensors",
|
| 205 |
+
"blocks.26.mlp.fc2.weight": "model-00007-of-00013.safetensors",
|
| 206 |
+
"blocks.26.norm1.weight": "model-00007-of-00013.safetensors",
|
| 207 |
+
"blocks.26.norm2.weight": "model-00007-of-00013.safetensors",
|
| 208 |
+
"blocks.27.adaln_proj.linear.bias": "model-00008-of-00013.safetensors",
|
| 209 |
+
"blocks.27.adaln_proj.linear.weight": "model-00008-of-00013.safetensors",
|
| 210 |
+
"blocks.27.attn.k_norm.weight": "model-00008-of-00013.safetensors",
|
| 211 |
+
"blocks.27.attn.out_proj.weight": "model-00008-of-00013.safetensors",
|
| 212 |
+
"blocks.27.attn.q_norm.weight": "model-00008-of-00013.safetensors",
|
| 213 |
+
"blocks.27.attn.qkv_proj.weight": "model-00008-of-00013.safetensors",
|
| 214 |
+
"blocks.27.mlp.fc1.weight": "model-00008-of-00013.safetensors",
|
| 215 |
+
"blocks.27.mlp.fc2.weight": "model-00008-of-00013.safetensors",
|
| 216 |
+
"blocks.27.norm1.weight": "model-00008-of-00013.safetensors",
|
| 217 |
+
"blocks.27.norm2.weight": "model-00008-of-00013.safetensors",
|
| 218 |
+
"blocks.28.adaln_proj.linear.bias": "model-00008-of-00013.safetensors",
|
| 219 |
+
"blocks.28.adaln_proj.linear.weight": "model-00008-of-00013.safetensors",
|
| 220 |
+
"blocks.28.attn.k_norm.weight": "model-00008-of-00013.safetensors",
|
| 221 |
+
"blocks.28.attn.out_proj.weight": "model-00008-of-00013.safetensors",
|
| 222 |
+
"blocks.28.attn.q_norm.weight": "model-00008-of-00013.safetensors",
|
| 223 |
+
"blocks.28.attn.qkv_proj.weight": "model-00008-of-00013.safetensors",
|
| 224 |
+
"blocks.28.mlp.fc1.weight": "model-00008-of-00013.safetensors",
|
| 225 |
+
"blocks.28.mlp.fc2.weight": "model-00008-of-00013.safetensors",
|
| 226 |
+
"blocks.28.norm1.weight": "model-00008-of-00013.safetensors",
|
| 227 |
+
"blocks.28.norm2.weight": "model-00008-of-00013.safetensors",
|
| 228 |
+
"blocks.29.adaln_proj.linear.bias": "model-00008-of-00013.safetensors",
|
| 229 |
+
"blocks.29.adaln_proj.linear.weight": "model-00008-of-00013.safetensors",
|
| 230 |
+
"blocks.29.attn.k_norm.weight": "model-00008-of-00013.safetensors",
|
| 231 |
+
"blocks.29.attn.out_proj.weight": "model-00008-of-00013.safetensors",
|
| 232 |
+
"blocks.29.attn.q_norm.weight": "model-00008-of-00013.safetensors",
|
| 233 |
+
"blocks.29.attn.qkv_proj.weight": "model-00008-of-00013.safetensors",
|
| 234 |
+
"blocks.29.mlp.fc1.weight": "model-00008-of-00013.safetensors",
|
| 235 |
+
"blocks.29.mlp.fc2.weight": "model-00008-of-00013.safetensors",
|
| 236 |
+
"blocks.29.norm1.weight": "model-00008-of-00013.safetensors",
|
| 237 |
+
"blocks.29.norm2.weight": "model-00008-of-00013.safetensors",
|
| 238 |
+
"blocks.3.adaln_proj.linear.bias": "model-00002-of-00013.safetensors",
|
| 239 |
+
"blocks.3.adaln_proj.linear.weight": "model-00002-of-00013.safetensors",
|
| 240 |
+
"blocks.3.attn.k_norm.weight": "model-00002-of-00013.safetensors",
|
| 241 |
+
"blocks.3.attn.out_proj.weight": "model-00002-of-00013.safetensors",
|
| 242 |
+
"blocks.3.attn.q_norm.weight": "model-00002-of-00013.safetensors",
|
| 243 |
+
"blocks.3.attn.qkv_proj.weight": "model-00002-of-00013.safetensors",
|
| 244 |
+
"blocks.3.mlp.fc1.weight": "model-00002-of-00013.safetensors",
|
| 245 |
+
"blocks.3.mlp.fc2.weight": "model-00002-of-00013.safetensors",
|
| 246 |
+
"blocks.3.norm1.weight": "model-00002-of-00013.safetensors",
|
| 247 |
+
"blocks.3.norm2.weight": "model-00002-of-00013.safetensors",
|
| 248 |
+
"blocks.30.adaln_proj.linear.bias": "model-00008-of-00013.safetensors",
|
| 249 |
+
"blocks.30.adaln_proj.linear.weight": "model-00008-of-00013.safetensors",
|
| 250 |
+
"blocks.30.attn.k_norm.weight": "model-00009-of-00013.safetensors",
|
| 251 |
+
"blocks.30.attn.out_proj.weight": "model-00009-of-00013.safetensors",
|
| 252 |
+
"blocks.30.attn.q_norm.weight": "model-00009-of-00013.safetensors",
|
| 253 |
+
"blocks.30.attn.qkv_proj.weight": "model-00009-of-00013.safetensors",
|
| 254 |
+
"blocks.30.mlp.fc1.weight": "model-00008-of-00013.safetensors",
|
| 255 |
+
"blocks.30.mlp.fc2.weight": "model-00008-of-00013.safetensors",
|
| 256 |
+
"blocks.30.norm1.weight": "model-00008-of-00013.safetensors",
|
| 257 |
+
"blocks.30.norm2.weight": "model-00008-of-00013.safetensors",
|
| 258 |
+
"blocks.31.adaln_proj.linear.bias": "model-00009-of-00013.safetensors",
|
| 259 |
+
"blocks.31.adaln_proj.linear.weight": "model-00009-of-00013.safetensors",
|
| 260 |
+
"blocks.31.attn.k_norm.weight": "model-00009-of-00013.safetensors",
|
| 261 |
+
"blocks.31.attn.out_proj.weight": "model-00009-of-00013.safetensors",
|
| 262 |
+
"blocks.31.attn.q_norm.weight": "model-00009-of-00013.safetensors",
|
| 263 |
+
"blocks.31.attn.qkv_proj.weight": "model-00009-of-00013.safetensors",
|
| 264 |
+
"blocks.31.mlp.fc1.weight": "model-00009-of-00013.safetensors",
|
| 265 |
+
"blocks.31.mlp.fc2.weight": "model-00009-of-00013.safetensors",
|
| 266 |
+
"blocks.31.norm1.weight": "model-00009-of-00013.safetensors",
|
| 267 |
+
"blocks.31.norm2.weight": "model-00009-of-00013.safetensors",
|
| 268 |
+
"blocks.32.adaln_proj.linear.bias": "model-00009-of-00013.safetensors",
|
| 269 |
+
"blocks.32.adaln_proj.linear.weight": "model-00009-of-00013.safetensors",
|
| 270 |
+
"blocks.32.attn.k_norm.weight": "model-00009-of-00013.safetensors",
|
| 271 |
+
"blocks.32.attn.out_proj.weight": "model-00009-of-00013.safetensors",
|
| 272 |
+
"blocks.32.attn.q_norm.weight": "model-00009-of-00013.safetensors",
|
| 273 |
+
"blocks.32.attn.qkv_proj.weight": "model-00009-of-00013.safetensors",
|
| 274 |
+
"blocks.32.mlp.fc1.weight": "model-00009-of-00013.safetensors",
|
| 275 |
+
"blocks.32.mlp.fc2.weight": "model-00009-of-00013.safetensors",
|
| 276 |
+
"blocks.32.norm1.weight": "model-00009-of-00013.safetensors",
|
| 277 |
+
"blocks.32.norm2.weight": "model-00009-of-00013.safetensors",
|
| 278 |
+
"blocks.33.adaln_proj.linear.bias": "model-00009-of-00013.safetensors",
|
| 279 |
+
"blocks.33.adaln_proj.linear.weight": "model-00009-of-00013.safetensors",
|
| 280 |
+
"blocks.33.attn.k_norm.weight": "model-00009-of-00013.safetensors",
|
| 281 |
+
"blocks.33.attn.out_proj.weight": "model-00009-of-00013.safetensors",
|
| 282 |
+
"blocks.33.attn.q_norm.weight": "model-00009-of-00013.safetensors",
|
| 283 |
+
"blocks.33.attn.qkv_proj.weight": "model-00009-of-00013.safetensors",
|
| 284 |
+
"blocks.33.mlp.fc1.weight": "model-00009-of-00013.safetensors",
|
| 285 |
+
"blocks.33.mlp.fc2.weight": "model-00009-of-00013.safetensors",
|
| 286 |
+
"blocks.33.norm1.weight": "model-00009-of-00013.safetensors",
|
| 287 |
+
"blocks.33.norm2.weight": "model-00009-of-00013.safetensors",
|
| 288 |
+
"blocks.34.adaln_proj.linear.bias": "model-00009-of-00013.safetensors",
|
| 289 |
+
"blocks.34.adaln_proj.linear.weight": "model-00009-of-00013.safetensors",
|
| 290 |
+
"blocks.34.attn.k_norm.weight": "model-00010-of-00013.safetensors",
|
| 291 |
+
"blocks.34.attn.out_proj.weight": "model-00010-of-00013.safetensors",
|
| 292 |
+
"blocks.34.attn.q_norm.weight": "model-00010-of-00013.safetensors",
|
| 293 |
+
"blocks.34.attn.qkv_proj.weight": "model-00010-of-00013.safetensors",
|
| 294 |
+
"blocks.34.mlp.fc1.weight": "model-00009-of-00013.safetensors",
|
| 295 |
+
"blocks.34.mlp.fc2.weight": "model-00009-of-00013.safetensors",
|
| 296 |
+
"blocks.34.norm1.weight": "model-00009-of-00013.safetensors",
|
| 297 |
+
"blocks.34.norm2.weight": "model-00009-of-00013.safetensors",
|
| 298 |
+
"blocks.35.adaln_proj.linear.bias": "model-00010-of-00013.safetensors",
|
| 299 |
+
"blocks.35.adaln_proj.linear.weight": "model-00010-of-00013.safetensors",
|
| 300 |
+
"blocks.35.attn.k_norm.weight": "model-00010-of-00013.safetensors",
|
| 301 |
+
"blocks.35.attn.out_proj.weight": "model-00010-of-00013.safetensors",
|
| 302 |
+
"blocks.35.attn.q_norm.weight": "model-00010-of-00013.safetensors",
|
| 303 |
+
"blocks.35.attn.qkv_proj.weight": "model-00010-of-00013.safetensors",
|
| 304 |
+
"blocks.35.mlp.fc1.weight": "model-00010-of-00013.safetensors",
|
| 305 |
+
"blocks.35.mlp.fc2.weight": "model-00010-of-00013.safetensors",
|
| 306 |
+
"blocks.35.norm1.weight": "model-00010-of-00013.safetensors",
|
| 307 |
+
"blocks.35.norm2.weight": "model-00010-of-00013.safetensors",
|
| 308 |
+
"blocks.36.adaln_proj.linear.bias": "model-00010-of-00013.safetensors",
|
| 309 |
+
"blocks.36.adaln_proj.linear.weight": "model-00010-of-00013.safetensors",
|
| 310 |
+
"blocks.36.attn.k_norm.weight": "model-00010-of-00013.safetensors",
|
| 311 |
+
"blocks.36.attn.out_proj.weight": "model-00010-of-00013.safetensors",
|
| 312 |
+
"blocks.36.attn.q_norm.weight": "model-00010-of-00013.safetensors",
|
| 313 |
+
"blocks.36.attn.qkv_proj.weight": "model-00010-of-00013.safetensors",
|
| 314 |
+
"blocks.36.mlp.fc1.weight": "model-00010-of-00013.safetensors",
|
| 315 |
+
"blocks.36.mlp.fc2.weight": "model-00010-of-00013.safetensors",
|
| 316 |
+
"blocks.36.norm1.weight": "model-00010-of-00013.safetensors",
|
| 317 |
+
"blocks.36.norm2.weight": "model-00010-of-00013.safetensors",
|
| 318 |
+
"blocks.37.adaln_proj.linear.bias": "model-00010-of-00013.safetensors",
|
| 319 |
+
"blocks.37.adaln_proj.linear.weight": "model-00010-of-00013.safetensors",
|
| 320 |
+
"blocks.37.attn.k_norm.weight": "model-00010-of-00013.safetensors",
|
| 321 |
+
"blocks.37.attn.out_proj.weight": "model-00010-of-00013.safetensors",
|
| 322 |
+
"blocks.37.attn.q_norm.weight": "model-00010-of-00013.safetensors",
|
| 323 |
+
"blocks.37.attn.qkv_proj.weight": "model-00010-of-00013.safetensors",
|
| 324 |
+
"blocks.37.mlp.fc1.weight": "model-00010-of-00013.safetensors",
|
| 325 |
+
"blocks.37.mlp.fc2.weight": "model-00010-of-00013.safetensors",
|
| 326 |
+
"blocks.37.norm1.weight": "model-00010-of-00013.safetensors",
|
| 327 |
+
"blocks.37.norm2.weight": "model-00010-of-00013.safetensors",
|
| 328 |
+
"blocks.38.adaln_proj.linear.bias": "model-00010-of-00013.safetensors",
|
| 329 |
+
"blocks.38.adaln_proj.linear.weight": "model-00010-of-00013.safetensors",
|
| 330 |
+
"blocks.38.attn.k_norm.weight": "model-00011-of-00013.safetensors",
|
| 331 |
+
"blocks.38.attn.out_proj.weight": "model-00011-of-00013.safetensors",
|
| 332 |
+
"blocks.38.attn.q_norm.weight": "model-00011-of-00013.safetensors",
|
| 333 |
+
"blocks.38.attn.qkv_proj.weight": "model-00011-of-00013.safetensors",
|
| 334 |
+
"blocks.38.mlp.fc1.weight": "model-00010-of-00013.safetensors",
|
| 335 |
+
"blocks.38.mlp.fc2.weight": "model-00010-of-00013.safetensors",
|
| 336 |
+
"blocks.38.norm1.weight": "model-00010-of-00013.safetensors",
|
| 337 |
+
"blocks.38.norm2.weight": "model-00010-of-00013.safetensors",
|
| 338 |
+
"blocks.39.adaln_proj.linear.bias": "model-00011-of-00013.safetensors",
|
| 339 |
+
"blocks.39.adaln_proj.linear.weight": "model-00011-of-00013.safetensors",
|
| 340 |
+
"blocks.39.attn.k_norm.weight": "model-00011-of-00013.safetensors",
|
| 341 |
+
"blocks.39.attn.out_proj.weight": "model-00011-of-00013.safetensors",
|
| 342 |
+
"blocks.39.attn.q_norm.weight": "model-00011-of-00013.safetensors",
|
| 343 |
+
"blocks.39.attn.qkv_proj.weight": "model-00011-of-00013.safetensors",
|
| 344 |
+
"blocks.39.mlp.fc1.weight": "model-00011-of-00013.safetensors",
|
| 345 |
+
"blocks.39.mlp.fc2.weight": "model-00011-of-00013.safetensors",
|
| 346 |
+
"blocks.39.norm1.weight": "model-00011-of-00013.safetensors",
|
| 347 |
+
"blocks.39.norm2.weight": "model-00011-of-00013.safetensors",
|
| 348 |
+
"blocks.4.adaln_proj.linear.bias": "model-00002-of-00013.safetensors",
|
| 349 |
+
"blocks.4.adaln_proj.linear.weight": "model-00002-of-00013.safetensors",
|
| 350 |
+
"blocks.4.attn.k_norm.weight": "model-00002-of-00013.safetensors",
|
| 351 |
+
"blocks.4.attn.out_proj.weight": "model-00002-of-00013.safetensors",
|
| 352 |
+
"blocks.4.attn.q_norm.weight": "model-00002-of-00013.safetensors",
|
| 353 |
+
"blocks.4.attn.qkv_proj.weight": "model-00002-of-00013.safetensors",
|
| 354 |
+
"blocks.4.mlp.fc1.weight": "model-00002-of-00013.safetensors",
|
| 355 |
+
"blocks.4.mlp.fc2.weight": "model-00002-of-00013.safetensors",
|
| 356 |
+
"blocks.4.norm1.weight": "model-00002-of-00013.safetensors",
|
| 357 |
+
"blocks.4.norm2.weight": "model-00002-of-00013.safetensors",
|
| 358 |
+
"blocks.40.adaln_proj.linear.bias": "model-00011-of-00013.safetensors",
|
| 359 |
+
"blocks.40.adaln_proj.linear.weight": "model-00011-of-00013.safetensors",
|
| 360 |
+
"blocks.40.attn.k_norm.weight": "model-00011-of-00013.safetensors",
|
| 361 |
+
"blocks.40.attn.out_proj.weight": "model-00011-of-00013.safetensors",
|
| 362 |
+
"blocks.40.attn.q_norm.weight": "model-00011-of-00013.safetensors",
|
| 363 |
+
"blocks.40.attn.qkv_proj.weight": "model-00011-of-00013.safetensors",
|
| 364 |
+
"blocks.40.mlp.fc1.weight": "model-00011-of-00013.safetensors",
|
| 365 |
+
"blocks.40.mlp.fc2.weight": "model-00011-of-00013.safetensors",
|
| 366 |
+
"blocks.40.norm1.weight": "model-00011-of-00013.safetensors",
|
| 367 |
+
"blocks.40.norm2.weight": "model-00011-of-00013.safetensors",
|
| 368 |
+
"blocks.41.adaln_proj.linear.bias": "model-00011-of-00013.safetensors",
|
| 369 |
+
"blocks.41.adaln_proj.linear.weight": "model-00011-of-00013.safetensors",
|
| 370 |
+
"blocks.41.attn.k_norm.weight": "model-00011-of-00013.safetensors",
|
| 371 |
+
"blocks.41.attn.out_proj.weight": "model-00011-of-00013.safetensors",
|
| 372 |
+
"blocks.41.attn.q_norm.weight": "model-00011-of-00013.safetensors",
|
| 373 |
+
"blocks.41.attn.qkv_proj.weight": "model-00011-of-00013.safetensors",
|
| 374 |
+
"blocks.41.mlp.fc1.weight": "model-00011-of-00013.safetensors",
|
| 375 |
+
"blocks.41.mlp.fc2.weight": "model-00011-of-00013.safetensors",
|
| 376 |
+
"blocks.41.norm1.weight": "model-00011-of-00013.safetensors",
|
| 377 |
+
"blocks.41.norm2.weight": "model-00011-of-00013.safetensors",
|
| 378 |
+
"blocks.42.adaln_proj.linear.bias": "model-00011-of-00013.safetensors",
|
| 379 |
+
"blocks.42.adaln_proj.linear.weight": "model-00011-of-00013.safetensors",
|
| 380 |
+
"blocks.42.attn.k_norm.weight": "model-00012-of-00013.safetensors",
|
| 381 |
+
"blocks.42.attn.out_proj.weight": "model-00012-of-00013.safetensors",
|
| 382 |
+
"blocks.42.attn.q_norm.weight": "model-00012-of-00013.safetensors",
|
| 383 |
+
"blocks.42.attn.qkv_proj.weight": "model-00012-of-00013.safetensors",
|
| 384 |
+
"blocks.42.mlp.fc1.weight": "model-00011-of-00013.safetensors",
|
| 385 |
+
"blocks.42.mlp.fc2.weight": "model-00011-of-00013.safetensors",
|
| 386 |
+
"blocks.42.norm1.weight": "model-00011-of-00013.safetensors",
|
| 387 |
+
"blocks.42.norm2.weight": "model-00011-of-00013.safetensors",
|
| 388 |
+
"blocks.43.adaln_proj.linear.bias": "model-00012-of-00013.safetensors",
|
| 389 |
+
"blocks.43.adaln_proj.linear.weight": "model-00012-of-00013.safetensors",
|
| 390 |
+
"blocks.43.attn.k_norm.weight": "model-00012-of-00013.safetensors",
|
| 391 |
+
"blocks.43.attn.out_proj.weight": "model-00012-of-00013.safetensors",
|
| 392 |
+
"blocks.43.attn.q_norm.weight": "model-00012-of-00013.safetensors",
|
| 393 |
+
"blocks.43.attn.qkv_proj.weight": "model-00012-of-00013.safetensors",
|
| 394 |
+
"blocks.43.mlp.fc1.weight": "model-00012-of-00013.safetensors",
|
| 395 |
+
"blocks.43.mlp.fc2.weight": "model-00012-of-00013.safetensors",
|
| 396 |
+
"blocks.43.norm1.weight": "model-00012-of-00013.safetensors",
|
| 397 |
+
"blocks.43.norm2.weight": "model-00012-of-00013.safetensors",
|
| 398 |
+
"blocks.44.adaln_proj.linear.bias": "model-00012-of-00013.safetensors",
|
| 399 |
+
"blocks.44.adaln_proj.linear.weight": "model-00012-of-00013.safetensors",
|
| 400 |
+
"blocks.44.attn.k_norm.weight": "model-00012-of-00013.safetensors",
|
| 401 |
+
"blocks.44.attn.out_proj.weight": "model-00012-of-00013.safetensors",
|
| 402 |
+
"blocks.44.attn.q_norm.weight": "model-00012-of-00013.safetensors",
|
| 403 |
+
"blocks.44.attn.qkv_proj.weight": "model-00012-of-00013.safetensors",
|
| 404 |
+
"blocks.44.mlp.fc1.weight": "model-00012-of-00013.safetensors",
|
| 405 |
+
"blocks.44.mlp.fc2.weight": "model-00012-of-00013.safetensors",
|
| 406 |
+
"blocks.44.norm1.weight": "model-00012-of-00013.safetensors",
|
| 407 |
+
"blocks.44.norm2.weight": "model-00012-of-00013.safetensors",
|
| 408 |
+
"blocks.45.adaln_proj.linear.bias": "model-00012-of-00013.safetensors",
|
| 409 |
+
"blocks.45.adaln_proj.linear.weight": "model-00012-of-00013.safetensors",
|
| 410 |
+
"blocks.45.attn.k_norm.weight": "model-00012-of-00013.safetensors",
|
| 411 |
+
"blocks.45.attn.out_proj.weight": "model-00012-of-00013.safetensors",
|
| 412 |
+
"blocks.45.attn.q_norm.weight": "model-00012-of-00013.safetensors",
|
| 413 |
+
"blocks.45.attn.qkv_proj.weight": "model-00012-of-00013.safetensors",
|
| 414 |
+
"blocks.45.mlp.fc1.weight": "model-00012-of-00013.safetensors",
|
| 415 |
+
"blocks.45.mlp.fc2.weight": "model-00012-of-00013.safetensors",
|
| 416 |
+
"blocks.45.norm1.weight": "model-00012-of-00013.safetensors",
|
| 417 |
+
"blocks.45.norm2.weight": "model-00012-of-00013.safetensors",
|
| 418 |
+
"blocks.46.adaln_proj.linear.bias": "model-00012-of-00013.safetensors",
|
| 419 |
+
"blocks.46.adaln_proj.linear.weight": "model-00012-of-00013.safetensors",
|
| 420 |
+
"blocks.46.attn.k_norm.weight": "model-00013-of-00013.safetensors",
|
| 421 |
+
"blocks.46.attn.out_proj.weight": "model-00013-of-00013.safetensors",
|
| 422 |
+
"blocks.46.attn.q_norm.weight": "model-00013-of-00013.safetensors",
|
| 423 |
+
"blocks.46.attn.qkv_proj.weight": "model-00013-of-00013.safetensors",
|
| 424 |
+
"blocks.46.mlp.fc1.weight": "model-00012-of-00013.safetensors",
|
| 425 |
+
"blocks.46.mlp.fc2.weight": "model-00012-of-00013.safetensors",
|
| 426 |
+
"blocks.46.norm1.weight": "model-00012-of-00013.safetensors",
|
| 427 |
+
"blocks.46.norm2.weight": "model-00012-of-00013.safetensors",
|
| 428 |
+
"blocks.47.adaln_proj.linear.bias": "model-00013-of-00013.safetensors",
|
| 429 |
+
"blocks.47.adaln_proj.linear.weight": "model-00013-of-00013.safetensors",
|
| 430 |
+
"blocks.47.attn.k_norm.weight": "model-00013-of-00013.safetensors",
|
| 431 |
+
"blocks.47.attn.out_proj.weight": "model-00013-of-00013.safetensors",
|
| 432 |
+
"blocks.47.attn.q_norm.weight": "model-00013-of-00013.safetensors",
|
| 433 |
+
"blocks.47.attn.qkv_proj.weight": "model-00013-of-00013.safetensors",
|
| 434 |
+
"blocks.47.mlp.fc1.weight": "model-00013-of-00013.safetensors",
|
| 435 |
+
"blocks.47.mlp.fc2.weight": "model-00013-of-00013.safetensors",
|
| 436 |
+
"blocks.47.norm1.weight": "model-00013-of-00013.safetensors",
|
| 437 |
+
"blocks.47.norm2.weight": "model-00013-of-00013.safetensors",
|
| 438 |
+
"blocks.48.adaln_proj.linear.bias": "model-00013-of-00013.safetensors",
|
| 439 |
+
"blocks.48.adaln_proj.linear.weight": "model-00013-of-00013.safetensors",
|
| 440 |
+
"blocks.48.attn.k_norm.weight": "model-00013-of-00013.safetensors",
|
| 441 |
+
"blocks.48.attn.out_proj.weight": "model-00013-of-00013.safetensors",
|
| 442 |
+
"blocks.48.attn.q_norm.weight": "model-00013-of-00013.safetensors",
|
| 443 |
+
"blocks.48.attn.qkv_proj.weight": "model-00013-of-00013.safetensors",
|
| 444 |
+
"blocks.48.mlp.fc1.weight": "model-00013-of-00013.safetensors",
|
| 445 |
+
"blocks.48.mlp.fc2.weight": "model-00013-of-00013.safetensors",
|
| 446 |
+
"blocks.48.norm1.weight": "model-00013-of-00013.safetensors",
|
| 447 |
+
"blocks.48.norm2.weight": "model-00013-of-00013.safetensors",
|
| 448 |
+
"blocks.49.adaln_proj.linear.bias": "model-00013-of-00013.safetensors",
|
| 449 |
+
"blocks.49.adaln_proj.linear.weight": "model-00013-of-00013.safetensors",
|
| 450 |
+
"blocks.49.attn.k_norm.weight": "model-00013-of-00013.safetensors",
|
| 451 |
+
"blocks.49.attn.out_proj.weight": "model-00013-of-00013.safetensors",
|
| 452 |
+
"blocks.49.attn.q_norm.weight": "model-00013-of-00013.safetensors",
|
| 453 |
+
"blocks.49.attn.qkv_proj.weight": "model-00013-of-00013.safetensors",
|
| 454 |
+
"blocks.49.mlp.fc1.weight": "model-00013-of-00013.safetensors",
|
| 455 |
+
"blocks.49.mlp.fc2.weight": "model-00013-of-00013.safetensors",
|
| 456 |
+
"blocks.49.norm1.weight": "model-00013-of-00013.safetensors",
|
| 457 |
+
"blocks.49.norm2.weight": "model-00013-of-00013.safetensors",
|
| 458 |
+
"blocks.5.adaln_proj.linear.bias": "model-00002-of-00013.safetensors",
|
| 459 |
+
"blocks.5.adaln_proj.linear.weight": "model-00002-of-00013.safetensors",
|
| 460 |
+
"blocks.5.attn.k_norm.weight": "model-00002-of-00013.safetensors",
|
| 461 |
+
"blocks.5.attn.out_proj.weight": "model-00002-of-00013.safetensors",
|
| 462 |
+
"blocks.5.attn.q_norm.weight": "model-00002-of-00013.safetensors",
|
| 463 |
+
"blocks.5.attn.qkv_proj.weight": "model-00002-of-00013.safetensors",
|
| 464 |
+
"blocks.5.mlp.fc1.weight": "model-00002-of-00013.safetensors",
|
| 465 |
+
"blocks.5.mlp.fc2.weight": "model-00002-of-00013.safetensors",
|
| 466 |
+
"blocks.5.norm1.weight": "model-00002-of-00013.safetensors",
|
| 467 |
+
"blocks.5.norm2.weight": "model-00002-of-00013.safetensors",
|
| 468 |
+
"blocks.6.adaln_proj.linear.bias": "model-00002-of-00013.safetensors",
|
| 469 |
+
"blocks.6.adaln_proj.linear.weight": "model-00002-of-00013.safetensors",
|
| 470 |
+
"blocks.6.attn.k_norm.weight": "model-00003-of-00013.safetensors",
|
| 471 |
+
"blocks.6.attn.out_proj.weight": "model-00003-of-00013.safetensors",
|
| 472 |
+
"blocks.6.attn.q_norm.weight": "model-00003-of-00013.safetensors",
|
| 473 |
+
"blocks.6.attn.qkv_proj.weight": "model-00003-of-00013.safetensors",
|
| 474 |
+
"blocks.6.mlp.fc1.weight": "model-00002-of-00013.safetensors",
|
| 475 |
+
"blocks.6.mlp.fc2.weight": "model-00002-of-00013.safetensors",
|
| 476 |
+
"blocks.6.norm1.weight": "model-00002-of-00013.safetensors",
|
| 477 |
+
"blocks.6.norm2.weight": "model-00002-of-00013.safetensors",
|
| 478 |
+
"blocks.7.adaln_proj.linear.bias": "model-00003-of-00013.safetensors",
|
| 479 |
+
"blocks.7.adaln_proj.linear.weight": "model-00003-of-00013.safetensors",
|
| 480 |
+
"blocks.7.attn.k_norm.weight": "model-00003-of-00013.safetensors",
|
| 481 |
+
"blocks.7.attn.out_proj.weight": "model-00003-of-00013.safetensors",
|
| 482 |
+
"blocks.7.attn.q_norm.weight": "model-00003-of-00013.safetensors",
|
| 483 |
+
"blocks.7.attn.qkv_proj.weight": "model-00003-of-00013.safetensors",
|
| 484 |
+
"blocks.7.mlp.fc1.weight": "model-00003-of-00013.safetensors",
|
| 485 |
+
"blocks.7.mlp.fc2.weight": "model-00003-of-00013.safetensors",
|
| 486 |
+
"blocks.7.norm1.weight": "model-00003-of-00013.safetensors",
|
| 487 |
+
"blocks.7.norm2.weight": "model-00003-of-00013.safetensors",
|
| 488 |
+
"blocks.8.adaln_proj.linear.bias": "model-00003-of-00013.safetensors",
|
| 489 |
+
"blocks.8.adaln_proj.linear.weight": "model-00003-of-00013.safetensors",
|
| 490 |
+
"blocks.8.attn.k_norm.weight": "model-00003-of-00013.safetensors",
|
| 491 |
+
"blocks.8.attn.out_proj.weight": "model-00003-of-00013.safetensors",
|
| 492 |
+
"blocks.8.attn.q_norm.weight": "model-00003-of-00013.safetensors",
|
| 493 |
+
"blocks.8.attn.qkv_proj.weight": "model-00003-of-00013.safetensors",
|
| 494 |
+
"blocks.8.mlp.fc1.weight": "model-00003-of-00013.safetensors",
|
| 495 |
+
"blocks.8.mlp.fc2.weight": "model-00003-of-00013.safetensors",
|
| 496 |
+
"blocks.8.norm1.weight": "model-00003-of-00013.safetensors",
|
| 497 |
+
"blocks.8.norm2.weight": "model-00003-of-00013.safetensors",
|
| 498 |
+
"blocks.9.adaln_proj.linear.bias": "model-00003-of-00013.safetensors",
|
| 499 |
+
"blocks.9.adaln_proj.linear.weight": "model-00003-of-00013.safetensors",
|
| 500 |
+
"blocks.9.attn.k_norm.weight": "model-00003-of-00013.safetensors",
|
| 501 |
+
"blocks.9.attn.out_proj.weight": "model-00003-of-00013.safetensors",
|
| 502 |
+
"blocks.9.attn.q_norm.weight": "model-00003-of-00013.safetensors",
|
| 503 |
+
"blocks.9.attn.qkv_proj.weight": "model-00003-of-00013.safetensors",
|
| 504 |
+
"blocks.9.mlp.fc1.weight": "model-00003-of-00013.safetensors",
|
| 505 |
+
"blocks.9.mlp.fc2.weight": "model-00003-of-00013.safetensors",
|
| 506 |
+
"blocks.9.norm1.weight": "model-00003-of-00013.safetensors",
|
| 507 |
+
"blocks.9.norm2.weight": "model-00003-of-00013.safetensors",
|
| 508 |
+
"condition_proj.bias": "model-00001-of-00013.safetensors",
|
| 509 |
+
"condition_proj.weight": "model-00001-of-00013.safetensors",
|
| 510 |
+
"final_layer.adaln_proj.linear.bias": "model-00013-of-00013.safetensors",
|
| 511 |
+
"final_layer.adaln_proj.linear.weight": "model-00013-of-00013.safetensors",
|
| 512 |
+
"final_layer.audio_out.bias": "model-00013-of-00013.safetensors",
|
| 513 |
+
"final_layer.audio_out.weight": "model-00013-of-00013.safetensors",
|
| 514 |
+
"final_layer.norm.weight": "model-00013-of-00013.safetensors",
|
| 515 |
+
"final_layer.video_out.bias": "model-00013-of-00013.safetensors",
|
| 516 |
+
"final_layer.video_out.weight": "model-00013-of-00013.safetensors",
|
| 517 |
+
"rope.inv_freq": "model-00001-of-00013.safetensors",
|
| 518 |
+
"time_embedder.proj_in.bias": "model-00001-of-00013.safetensors",
|
| 519 |
+
"time_embedder.proj_in.weight": "model-00001-of-00013.safetensors",
|
| 520 |
+
"time_embedder.proj_out.bias": "model-00001-of-00013.safetensors",
|
| 521 |
+
"time_embedder.proj_out.weight": "model-00001-of-00013.safetensors",
|
| 522 |
+
"token_refiner.blocks.0.attn.k_norm.weight": "model-00001-of-00013.safetensors",
|
| 523 |
+
"token_refiner.blocks.0.attn.out_proj.weight": "model-00001-of-00013.safetensors",
|
| 524 |
+
"token_refiner.blocks.0.attn.q_norm.weight": "model-00001-of-00013.safetensors",
|
| 525 |
+
"token_refiner.blocks.0.attn.qkv_proj.weight": "model-00001-of-00013.safetensors",
|
| 526 |
+
"token_refiner.blocks.0.mlp.fc1.weight": "model-00001-of-00013.safetensors",
|
| 527 |
+
"token_refiner.blocks.0.mlp.fc2.weight": "model-00001-of-00013.safetensors",
|
| 528 |
+
"token_refiner.blocks.0.norm1.weight": "model-00001-of-00013.safetensors",
|
| 529 |
+
"token_refiner.blocks.0.norm2.weight": "model-00001-of-00013.safetensors",
|
| 530 |
+
"token_refiner.blocks.1.attn.k_norm.weight": "model-00001-of-00013.safetensors",
|
| 531 |
+
"token_refiner.blocks.1.attn.out_proj.weight": "model-00001-of-00013.safetensors",
|
| 532 |
+
"token_refiner.blocks.1.attn.q_norm.weight": "model-00001-of-00013.safetensors",
|
| 533 |
+
"token_refiner.blocks.1.attn.qkv_proj.weight": "model-00001-of-00013.safetensors",
|
| 534 |
+
"token_refiner.blocks.1.mlp.fc1.weight": "model-00001-of-00013.safetensors",
|
| 535 |
+
"token_refiner.blocks.1.mlp.fc2.weight": "model-00001-of-00013.safetensors",
|
| 536 |
+
"token_refiner.blocks.1.norm1.weight": "model-00001-of-00013.safetensors",
|
| 537 |
+
"token_refiner.blocks.1.norm2.weight": "model-00001-of-00013.safetensors",
|
| 538 |
+
"token_refiner.final_norm.weight": "model-00001-of-00013.safetensors",
|
| 539 |
+
"video_patch_proj.bias": "model-00001-of-00013.safetensors",
|
| 540 |
+
"video_patch_proj.weight": "model-00001-of-00013.safetensors"
|
| 541 |
+
}
|
| 542 |
+
}
|
FL2VA/video_vae/attention.py
ADDED
|
@@ -0,0 +1,163 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 2 |
+
# Attention module for the MiniMax H3 visual VAE (inference-only bundle).
|
| 3 |
+
import os
|
| 4 |
+
import torch
|
| 5 |
+
import torch.nn as nn
|
| 6 |
+
import torch.distributed as dist
|
| 7 |
+
from typing import Optional
|
| 8 |
+
from diffusers.utils import logging
|
| 9 |
+
|
| 10 |
+
from .parallel import all_to_all_4D, get_parallel_state
|
| 11 |
+
from .func import apply_rotary_pos_emb
|
| 12 |
+
from .flash import flash_attn
|
| 13 |
+
|
| 14 |
+
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
| 15 |
+
|
| 16 |
+
|
| 17 |
+
def _env_flag(name, default="0"):
|
| 18 |
+
value = os.environ.get(name, default)
|
| 19 |
+
return str(value).strip().lower() in ("1", "true", "yes", "on")
|
| 20 |
+
|
| 21 |
+
|
| 22 |
+
def _vit_norm_input(module, hidden_states):
|
| 23 |
+
if _env_flag("MINIMAX_H3_VAE_DECODER_VIT_FP32_NORM", "1"):
|
| 24 |
+
return hidden_states.float()
|
| 25 |
+
weight = getattr(module, "weight", None)
|
| 26 |
+
return hidden_states.to(getattr(weight, "dtype", hidden_states.dtype))
|
| 27 |
+
|
| 28 |
+
|
| 29 |
+
def maybe_checkpoint(owner, function, *args):
|
| 30 |
+
if owner.training and getattr(owner, "gradient_checkpointing", False):
|
| 31 |
+
raise NotImplementedError(
|
| 32 |
+
"gradient checkpointing is not supported in this inference-only bundle"
|
| 33 |
+
)
|
| 34 |
+
return function(*args)
|
| 35 |
+
|
| 36 |
+
|
| 37 |
+
class Attention(nn.Module):
|
| 38 |
+
def __init__(
|
| 39 |
+
self,
|
| 40 |
+
heads,
|
| 41 |
+
dim_head,
|
| 42 |
+
embed_dim: Optional[int] = None,
|
| 43 |
+
qk_norm_type: Optional[str] = None,
|
| 44 |
+
qk_norm_affine: bool = False,
|
| 45 |
+
bias: bool = True,
|
| 46 |
+
out_bias: Optional[bool] = None,
|
| 47 |
+
eps: float = 1e-5,
|
| 48 |
+
**kwargs,
|
| 49 |
+
):
|
| 50 |
+
super().__init__()
|
| 51 |
+
self.dim_head = dim_head
|
| 52 |
+
self.heads = heads
|
| 53 |
+
self.attn_inner_dim = dim_head * heads
|
| 54 |
+
self.embed_dim = embed_dim if embed_dim is not None else self.attn_inner_dim
|
| 55 |
+
|
| 56 |
+
out_bias = out_bias if out_bias is not None else bias
|
| 57 |
+
|
| 58 |
+
if qk_norm_type is None:
|
| 59 |
+
self.norm_q = None
|
| 60 |
+
self.norm_k = None
|
| 61 |
+
elif qk_norm_type == "layer_norm":
|
| 62 |
+
self.norm_q = nn.LayerNorm(
|
| 63 |
+
dim_head, eps=eps, elementwise_affine=qk_norm_affine
|
| 64 |
+
)
|
| 65 |
+
self.norm_k = nn.LayerNorm(
|
| 66 |
+
dim_head, eps=eps, elementwise_affine=qk_norm_affine
|
| 67 |
+
)
|
| 68 |
+
elif qk_norm_type == "rms_norm":
|
| 69 |
+
self.norm_q = nn.RMSNorm(
|
| 70 |
+
dim_head, eps=eps, elementwise_affine=qk_norm_affine
|
| 71 |
+
)
|
| 72 |
+
self.norm_k = nn.RMSNorm(
|
| 73 |
+
dim_head, eps=eps, elementwise_affine=qk_norm_affine
|
| 74 |
+
)
|
| 75 |
+
else:
|
| 76 |
+
raise ValueError(
|
| 77 |
+
f"unknown qk_norm_type: {qk_norm_type}. Should be None,'layer_norm','rms_norm'"
|
| 78 |
+
)
|
| 79 |
+
|
| 80 |
+
self.to_qkv = nn.Linear(self.embed_dim, self.attn_inner_dim * 3, bias=bias)
|
| 81 |
+
|
| 82 |
+
self.to_out = nn.Linear(self.attn_inner_dim, self.embed_dim, bias=out_bias)
|
| 83 |
+
|
| 84 |
+
self.spatial_parallel = get_parallel_state().get("sp_enabled", False)
|
| 85 |
+
|
| 86 |
+
state = get_parallel_state()
|
| 87 |
+
sp_size = state.get("sp_size", 1)
|
| 88 |
+
tp_size = state.get("tp_size", 1)
|
| 89 |
+
parallel_size = sp_size * tp_size
|
| 90 |
+
if parallel_size > 1 and self.heads % parallel_size != 0:
|
| 91 |
+
raise ValueError(
|
| 92 |
+
f"num_heads {self.heads} must be divisible by sp_size * tp_size ({sp_size} * {tp_size} = {parallel_size})"
|
| 93 |
+
)
|
| 94 |
+
|
| 95 |
+
if len(kwargs) > 0 and (not dist.is_initialized() or dist.get_rank() == 0):
|
| 96 |
+
logger.warning(f"Unused kwargs: {kwargs}")
|
| 97 |
+
|
| 98 |
+
def _perform_attention(self, query, key, value, pack_info):
|
| 99 |
+
cu_seqlens = pack_info.get("cu_seqlens", None)
|
| 100 |
+
mask_mod = pack_info.get("mask_mod", None)
|
| 101 |
+
block_sparse = pack_info.get("block_sparse", None)
|
| 102 |
+
|
| 103 |
+
if cu_seqlens is not None:
|
| 104 |
+
raise NotImplementedError(
|
| 105 |
+
"varlen attention is not supported in this inference-only bundle"
|
| 106 |
+
)
|
| 107 |
+
|
| 108 |
+
if mask_mod is not None:
|
| 109 |
+
hidden_states = flash_attn(
|
| 110 |
+
query,
|
| 111 |
+
key,
|
| 112 |
+
value,
|
| 113 |
+
mask_mod=mask_mod,
|
| 114 |
+
block_sparse=block_sparse,
|
| 115 |
+
)
|
| 116 |
+
else:
|
| 117 |
+
hidden_states = flash_attn(
|
| 118 |
+
query,
|
| 119 |
+
key,
|
| 120 |
+
value,
|
| 121 |
+
)
|
| 122 |
+
|
| 123 |
+
return hidden_states
|
| 124 |
+
|
| 125 |
+
def perform_attention(self, query, key, value, pack_info={}):
|
| 126 |
+
return self._perform_attention(query, key, value, pack_info)
|
| 127 |
+
|
| 128 |
+
def forward(
|
| 129 |
+
self,
|
| 130 |
+
hidden_states: torch.Tensor,
|
| 131 |
+
rotary_pos_emb: Optional[torch.Tensor] = None,
|
| 132 |
+
pack_info: dict = {},
|
| 133 |
+
) -> torch.Tensor:
|
| 134 |
+
batch_size, seq_len, _ = hidden_states.shape
|
| 135 |
+
|
| 136 |
+
qkv = self.to_qkv(hidden_states)
|
| 137 |
+
qkv = qkv.view(batch_size, seq_len, -1, 3 * self.dim_head)
|
| 138 |
+
query, key, value = torch.chunk(qkv, 3, dim=-1)
|
| 139 |
+
|
| 140 |
+
if self.spatial_parallel:
|
| 141 |
+
local_process_group = get_parallel_state()["sp_process_group"]
|
| 142 |
+
query = all_to_all_4D(query, 2, 1, group=local_process_group)
|
| 143 |
+
key = all_to_all_4D(key, 2, 1, group=local_process_group)
|
| 144 |
+
value = all_to_all_4D(value, 2, 1, group=local_process_group)
|
| 145 |
+
|
| 146 |
+
if self.norm_q is not None:
|
| 147 |
+
query = self.norm_q(_vit_norm_input(self.norm_q, query)).to(query.dtype)
|
| 148 |
+
if self.norm_k is not None:
|
| 149 |
+
key = self.norm_k(_vit_norm_input(self.norm_k, key)).to(key.dtype)
|
| 150 |
+
|
| 151 |
+
if rotary_pos_emb is not None:
|
| 152 |
+
query = apply_rotary_pos_emb(query, rotary_pos_emb)
|
| 153 |
+
key = apply_rotary_pos_emb(key, rotary_pos_emb)
|
| 154 |
+
|
| 155 |
+
hidden_states = self.perform_attention(query, key, value, pack_info)
|
| 156 |
+
|
| 157 |
+
if self.spatial_parallel:
|
| 158 |
+
hidden_states = all_to_all_4D(hidden_states, 1, 2, group=local_process_group)
|
| 159 |
+
|
| 160 |
+
hidden_states = hidden_states.reshape(batch_size, seq_len, -1)
|
| 161 |
+
hidden_states = self.to_out(hidden_states)
|
| 162 |
+
|
| 163 |
+
return hidden_states
|
FL2VA/video_vae/base_module.py
ADDED
|
@@ -0,0 +1,282 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 2 |
+
# Transformer building blocks for the MiniMax H3 visual VAE ViT decoder.
|
| 3 |
+
import math
|
| 4 |
+
import os
|
| 5 |
+
import torch
|
| 6 |
+
import torch.nn as nn
|
| 7 |
+
from typing import Optional
|
| 8 |
+
from diffusers.utils import logging
|
| 9 |
+
from diffusers.utils.torch_utils import maybe_allow_in_graph
|
| 10 |
+
|
| 11 |
+
from .attention import Attention
|
| 12 |
+
|
| 13 |
+
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
| 14 |
+
|
| 15 |
+
|
| 16 |
+
def _env_flag(name, default="0"):
|
| 17 |
+
value = os.environ.get(name, default)
|
| 18 |
+
return str(value).strip().lower() in ("1", "true", "yes", "on")
|
| 19 |
+
|
| 20 |
+
|
| 21 |
+
def _env_optional_bool(name, default=""):
|
| 22 |
+
value = str(os.environ.get(name, default)).strip().lower()
|
| 23 |
+
if value in ("", "default", "auto", "none", "unset"):
|
| 24 |
+
return None
|
| 25 |
+
return value not in ("0", "false", "no", "off", "disabled")
|
| 26 |
+
|
| 27 |
+
|
| 28 |
+
def _vit_torch_compile_kwargs(prefix):
|
| 29 |
+
kwargs = {}
|
| 30 |
+
backend = os.environ.get(f"{prefix}_BACKEND", "inductor").strip()
|
| 31 |
+
mode = os.environ.get(f"{prefix}_MODE", "reduce-overhead").strip()
|
| 32 |
+
if backend and backend.lower() not in ("default", "none"):
|
| 33 |
+
kwargs["backend"] = backend
|
| 34 |
+
if mode and mode.lower() not in ("default", "none"):
|
| 35 |
+
kwargs["mode"] = mode
|
| 36 |
+
kwargs["fullgraph"] = _env_flag(f"{prefix}_FULLGRAPH", "0")
|
| 37 |
+
dynamic = _env_optional_bool(f"{prefix}_DYNAMIC")
|
| 38 |
+
if dynamic is not None:
|
| 39 |
+
kwargs["dynamic"] = dynamic
|
| 40 |
+
return kwargs
|
| 41 |
+
|
| 42 |
+
|
| 43 |
+
|
| 44 |
+
|
| 45 |
+
def _vit_norm_input(module, hidden_states):
|
| 46 |
+
if _env_flag("MINIMAX_H3_VAE_DECODER_VIT_FP32_NORM", "1"):
|
| 47 |
+
return hidden_states.float()
|
| 48 |
+
return hidden_states.to(getattr(module.weight, "dtype", hidden_states.dtype))
|
| 49 |
+
|
| 50 |
+
|
| 51 |
+
|
| 52 |
+
|
| 53 |
+
|
| 54 |
+
|
| 55 |
+
class FeedForward(nn.Module):
|
| 56 |
+
def __init__(
|
| 57 |
+
self,
|
| 58 |
+
dim: int,
|
| 59 |
+
dim_out: Optional[int] = None,
|
| 60 |
+
mult: int = 4,
|
| 61 |
+
activation_fn: str = "silu",
|
| 62 |
+
bias: bool = True,
|
| 63 |
+
use_gated: bool = True,
|
| 64 |
+
glu_balanced: bool = False,
|
| 65 |
+
):
|
| 66 |
+
super().__init__()
|
| 67 |
+
ratio = 2 / 3 if (use_gated and glu_balanced) else 1
|
| 68 |
+
inner_dim = round(dim * mult * ratio)
|
| 69 |
+
dim_out = dim_out if dim_out is not None else dim
|
| 70 |
+
self.use_gated = use_gated
|
| 71 |
+
|
| 72 |
+
if use_gated:
|
| 73 |
+
self.w1 = nn.Linear(dim, inner_dim * 2, bias=bias)
|
| 74 |
+
else:
|
| 75 |
+
self.w1 = nn.Linear(dim, inner_dim, bias=bias)
|
| 76 |
+
|
| 77 |
+
if activation_fn == "silu":
|
| 78 |
+
self.act_fn = nn.SiLU()
|
| 79 |
+
elif activation_fn == "gelu":
|
| 80 |
+
self.act_fn = nn.GELU()
|
| 81 |
+
elif activation_fn == "gelu-approximate":
|
| 82 |
+
self.act_fn = nn.GELU(approximate="tanh")
|
| 83 |
+
else:
|
| 84 |
+
raise ValueError(f"Unsupported activation function: {activation_fn}")
|
| 85 |
+
|
| 86 |
+
self.w2 = nn.Linear(inner_dim, dim_out, bias=bias)
|
| 87 |
+
self._compile_forward_enabled = _env_flag(
|
| 88 |
+
"MINIMAX_H3_VAE_DECODER_VIT_FF_TORCH_COMPILE", "0"
|
| 89 |
+
)
|
| 90 |
+
self._compile_forward_fatal = _env_flag(
|
| 91 |
+
"MINIMAX_H3_VAE_DECODER_VIT_FF_TORCH_COMPILE_FATAL", "0"
|
| 92 |
+
)
|
| 93 |
+
self._compiled_forward = None
|
| 94 |
+
|
| 95 |
+
def _forward_impl(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
| 96 |
+
hidden_states = self.w1(hidden_states)
|
| 97 |
+
|
| 98 |
+
if self.use_gated:
|
| 99 |
+
gate, hidden_states = hidden_states.chunk(2, dim=-1)
|
| 100 |
+
hidden_states = self.act_fn(gate) * hidden_states
|
| 101 |
+
else:
|
| 102 |
+
hidden_states = self.act_fn(hidden_states)
|
| 103 |
+
|
| 104 |
+
hidden_states = self.w2(hidden_states)
|
| 105 |
+
return hidden_states
|
| 106 |
+
|
| 107 |
+
def _get_forward_impl(self):
|
| 108 |
+
if not self._compile_forward_enabled:
|
| 109 |
+
return self._forward_impl
|
| 110 |
+
if self._compiled_forward is not None:
|
| 111 |
+
return self._compiled_forward
|
| 112 |
+
if not hasattr(torch, "compile"):
|
| 113 |
+
message = "torch.compile is unavailable; falling back to eager ViT FeedForward"
|
| 114 |
+
if self._compile_forward_fatal:
|
| 115 |
+
raise RuntimeError(message)
|
| 116 |
+
logger.warning(f"[ViTFeedForward] {message}")
|
| 117 |
+
self._compile_forward_enabled = False
|
| 118 |
+
return self._forward_impl
|
| 119 |
+
|
| 120 |
+
kwargs = _vit_torch_compile_kwargs("MINIMAX_H3_VAE_DECODER_VIT_FF_TORCH_COMPILE")
|
| 121 |
+
try:
|
| 122 |
+
self._compiled_forward = torch.compile(self._forward_impl, **kwargs)
|
| 123 |
+
logger.info(f"[ViTFeedForward] torch.compile enabled kwargs={kwargs}")
|
| 124 |
+
except Exception as exc:
|
| 125 |
+
if self._compile_forward_fatal:
|
| 126 |
+
raise
|
| 127 |
+
logger.warning(
|
| 128 |
+
f"[ViTFeedForward] torch.compile setup failed: {type(exc).__name__}: {exc}; "
|
| 129 |
+
"falling back to eager"
|
| 130 |
+
)
|
| 131 |
+
self._compile_forward_enabled = False
|
| 132 |
+
self._compiled_forward = None
|
| 133 |
+
return self._forward_impl
|
| 134 |
+
return self._compiled_forward
|
| 135 |
+
|
| 136 |
+
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
| 137 |
+
forward_impl = self._get_forward_impl()
|
| 138 |
+
try:
|
| 139 |
+
return forward_impl(hidden_states)
|
| 140 |
+
except Exception as exc:
|
| 141 |
+
if (
|
| 142 |
+
self._compile_forward_enabled
|
| 143 |
+
and self._compiled_forward is not None
|
| 144 |
+
and forward_impl is self._compiled_forward
|
| 145 |
+
and not self._compile_forward_fatal
|
| 146 |
+
):
|
| 147 |
+
logger.warning(
|
| 148 |
+
f"[ViTFeedForward] compiled forward failed: {type(exc).__name__}: {exc}; "
|
| 149 |
+
"disabling compile and retrying eager"
|
| 150 |
+
)
|
| 151 |
+
self._compile_forward_enabled = False
|
| 152 |
+
self._compiled_forward = None
|
| 153 |
+
return self._forward_impl(hidden_states)
|
| 154 |
+
raise
|
| 155 |
+
|
| 156 |
+
|
| 157 |
+
class RotaryEmbeddingND(nn.Module):
|
| 158 |
+
def __init__(self, dim, rotary_base=10000, n_dim=3, use_angle=False):
|
| 159 |
+
super().__init__()
|
| 160 |
+
self.dim = dim
|
| 161 |
+
self.n_dim = n_dim
|
| 162 |
+
|
| 163 |
+
if dim % (2 * n_dim) != 0:
|
| 164 |
+
raise ValueError(
|
| 165 |
+
f"head_dim {dim} must be divisible by 2 * n_dim {2 * n_dim}"
|
| 166 |
+
)
|
| 167 |
+
|
| 168 |
+
if use_angle:
|
| 169 |
+
self.angle_scale = 2.0 * math.pi
|
| 170 |
+
else:
|
| 171 |
+
self.angle_scale = 1.0
|
| 172 |
+
|
| 173 |
+
inv_freq = 1 / rotary_base ** torch.arange(
|
| 174 |
+
0, 1, 2 * n_dim / dim, dtype=torch.float32
|
| 175 |
+
)
|
| 176 |
+
self.register_buffer("inv_freq", inv_freq, persistent=False)
|
| 177 |
+
|
| 178 |
+
def forward(self, img_ids):
|
| 179 |
+
B, N, D = img_ids.shape
|
| 180 |
+
if D != self.n_dim:
|
| 181 |
+
raise ValueError(f"Expected {self.n_dim} dimensions, got {D}")
|
| 182 |
+
|
| 183 |
+
with torch.autocast("cuda", enabled=False):
|
| 184 |
+
angles = (
|
| 185 |
+
self.angle_scale
|
| 186 |
+
* img_ids[:, :, :, None]
|
| 187 |
+
* self.inv_freq.to(img_ids.device)[None, None, None, :]
|
| 188 |
+
)
|
| 189 |
+
angles = angles.flatten(2, 3)
|
| 190 |
+
angles = angles.tile(2)
|
| 191 |
+
angles = angles.unsqueeze(2)
|
| 192 |
+
|
| 193 |
+
cos = torch.cos(angles)
|
| 194 |
+
sin = torch.sin(angles)
|
| 195 |
+
|
| 196 |
+
return cos.to(dtype=img_ids.dtype), sin.to(dtype=img_ids.dtype)
|
| 197 |
+
|
| 198 |
+
|
| 199 |
+
@maybe_allow_in_graph
|
| 200 |
+
class TransformerBlock(nn.Module):
|
| 201 |
+
def __init__(
|
| 202 |
+
self,
|
| 203 |
+
heads: int,
|
| 204 |
+
dim_head: int,
|
| 205 |
+
embed_dim: Optional[int] = None,
|
| 206 |
+
ffn_glu_balanced: bool = False,
|
| 207 |
+
norm_type: str = "layer_norm",
|
| 208 |
+
norm_affine: bool = True,
|
| 209 |
+
qk_norm_type: str = "rms_norm",
|
| 210 |
+
qk_norm_affine: bool = False,
|
| 211 |
+
ffn_activation_fn: str = "silu",
|
| 212 |
+
ffn_use_gated: bool = True,
|
| 213 |
+
use_scale: bool = True,
|
| 214 |
+
bias: bool = True,
|
| 215 |
+
eps: float = 1e-5,
|
| 216 |
+
**kwargs,
|
| 217 |
+
):
|
| 218 |
+
super().__init__()
|
| 219 |
+
dim = embed_dim if embed_dim is not None else dim_head * heads
|
| 220 |
+
self.use_scale = use_scale
|
| 221 |
+
|
| 222 |
+
if norm_type == "layer_norm":
|
| 223 |
+
norm_class = nn.LayerNorm
|
| 224 |
+
elif norm_type == "rms_norm":
|
| 225 |
+
norm_class = nn.RMSNorm
|
| 226 |
+
else:
|
| 227 |
+
raise ValueError(f"unknown norm_type {norm_type}")
|
| 228 |
+
|
| 229 |
+
self.norm1 = norm_class(
|
| 230 |
+
dim,
|
| 231 |
+
elementwise_affine=norm_affine,
|
| 232 |
+
eps=eps,
|
| 233 |
+
)
|
| 234 |
+
self.attn = Attention(
|
| 235 |
+
heads=heads,
|
| 236 |
+
dim_head=dim_head,
|
| 237 |
+
embed_dim=dim,
|
| 238 |
+
qk_norm_type=qk_norm_type,
|
| 239 |
+
qk_norm_affine=qk_norm_affine,
|
| 240 |
+
bias=bias,
|
| 241 |
+
eps=eps,
|
| 242 |
+
**kwargs,
|
| 243 |
+
)
|
| 244 |
+
if use_scale:
|
| 245 |
+
self.scale1 = nn.Parameter(torch.zeros(dim))
|
| 246 |
+
|
| 247 |
+
self.norm2 = norm_class(
|
| 248 |
+
dim,
|
| 249 |
+
elementwise_affine=norm_affine,
|
| 250 |
+
eps=eps,
|
| 251 |
+
)
|
| 252 |
+
self.ff = FeedForward(
|
| 253 |
+
dim=dim,
|
| 254 |
+
activation_fn=ffn_activation_fn,
|
| 255 |
+
bias=bias,
|
| 256 |
+
use_gated=ffn_use_gated,
|
| 257 |
+
glu_balanced=ffn_glu_balanced,
|
| 258 |
+
)
|
| 259 |
+
if use_scale:
|
| 260 |
+
self.scale2 = nn.Parameter(torch.zeros(dim))
|
| 261 |
+
|
| 262 |
+
def forward(
|
| 263 |
+
self,
|
| 264 |
+
hidden_states: torch.FloatTensor,
|
| 265 |
+
rotary_pos_emb: Optional[torch.FloatTensor] = None,
|
| 266 |
+
pack_info: dict = {},
|
| 267 |
+
):
|
| 268 |
+
norm_hidden_states = self.norm1(_vit_norm_input(self.norm1, hidden_states)).to(hidden_states.dtype)
|
| 269 |
+
attn_output = self.attn(norm_hidden_states, rotary_pos_emb, pack_info)
|
| 270 |
+
if self.use_scale:
|
| 271 |
+
hidden_states = hidden_states + attn_output * self.scale1
|
| 272 |
+
else:
|
| 273 |
+
hidden_states = hidden_states + attn_output
|
| 274 |
+
|
| 275 |
+
norm_hidden_states = self.norm2(_vit_norm_input(self.norm2, hidden_states)).to(hidden_states.dtype)
|
| 276 |
+
ff_output = self.ff(norm_hidden_states)
|
| 277 |
+
if self.use_scale:
|
| 278 |
+
hidden_states = hidden_states + ff_output * self.scale2
|
| 279 |
+
else:
|
| 280 |
+
hidden_states = hidden_states + ff_output
|
| 281 |
+
|
| 282 |
+
return hidden_states
|
FL2VA/video_vae/config.json
ADDED
|
@@ -0,0 +1,74 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"_class_name": "MiniMaxH3VideoVAE",
|
| 3 |
+
"_diffusers_version": "0.32.2",
|
| 4 |
+
"mode": "standalone",
|
| 5 |
+
"auto_map": {
|
| 6 |
+
"AutoModel": "minimax_h3_video_vae.MiniMaxH3VideoVAE"
|
| 7 |
+
},
|
| 8 |
+
"source_path": "source",
|
| 9 |
+
"source_class_name": "AutoencoderKLLegacy",
|
| 10 |
+
"vae_clip_length": 17,
|
| 11 |
+
"vae_token_drop": 3,
|
| 12 |
+
"vae_encoder_tiling": 1,
|
| 13 |
+
"vae_decoder_tiling": 1,
|
| 14 |
+
"vae_parallel_tiling": 1,
|
| 15 |
+
"vae_tile_size": 256,
|
| 16 |
+
"vae_tile_overlap_min": 64,
|
| 17 |
+
"vae_encoder_parallel": 0,
|
| 18 |
+
"vae_decoder_parallel": 0,
|
| 19 |
+
"vae_chunk_dim": -1,
|
| 20 |
+
"source_safetensors_path": "model.safetensors",
|
| 21 |
+
"latent_channels": 24,
|
| 22 |
+
"latents_mean": [
|
| 23 |
+
0.858090341091156,
|
| 24 |
+
-0.9606591463088989,
|
| 25 |
+
1.0661640167236328,
|
| 26 |
+
-0.5090325474739075,
|
| 27 |
+
-0.2727581858634949,
|
| 28 |
+
-1.3675414323806763,
|
| 29 |
+
-0.2553254961967468,
|
| 30 |
+
-0.26907554268836975,
|
| 31 |
+
-0.5376840829849243,
|
| 32 |
+
-0.0464097298681736,
|
| 33 |
+
0.6657370328903198,
|
| 34 |
+
0.19690127670764923,
|
| 35 |
+
-0.5460608005523682,
|
| 36 |
+
-0.4035342037677765,
|
| 37 |
+
-0.23683024942874908,
|
| 38 |
+
0.25928452610969543,
|
| 39 |
+
-0.30133944749832153,
|
| 40 |
+
0.211341992020607,
|
| 41 |
+
-1.1206848621368408,
|
| 42 |
+
0.3581933379173279,
|
| 43 |
+
-0.04225143790245056,
|
| 44 |
+
0.2604829967021942,
|
| 45 |
+
0.22864092886447906,
|
| 46 |
+
0.7056031823158264
|
| 47 |
+
],
|
| 48 |
+
"latents_std": [
|
| 49 |
+
1.2223774194717407,
|
| 50 |
+
1.2767263650894165,
|
| 51 |
+
1.68317747116088865,
|
| 52 |
+
1.7549455165863037,
|
| 53 |
+
1.5636216402053833,
|
| 54 |
+
2.194143533706665,
|
| 55 |
+
0.96531379222869875,
|
| 56 |
+
1.05698859691619875,
|
| 57 |
+
0.841948926448822,
|
| 58 |
+
0.7729952931404114,
|
| 59 |
+
1.8955937623977661,
|
| 60 |
+
0.946841835975647,
|
| 61 |
+
0.7996809482574463,
|
| 62 |
+
0.44988900423049925,
|
| 63 |
+
0.7197399735450745,
|
| 64 |
+
0.69362932443618775,
|
| 65 |
+
2.961095094680786,
|
| 66 |
+
2.7694199085235595,
|
| 67 |
+
3.0496184825897215,
|
| 68 |
+
2.1088054180145265,
|
| 69 |
+
3.276226282119751,
|
| 70 |
+
3.1627357006073,
|
| 71 |
+
2.28168129920959475,
|
| 72 |
+
2.6127843856811525
|
| 73 |
+
]
|
| 74 |
+
}
|
FL2VA/video_vae/conv.py
ADDED
|
@@ -0,0 +1,159 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 2 |
+
# Spatial-parallel 3D convolution for the MiniMax H3 visual VAE.
|
| 3 |
+
import torch
|
| 4 |
+
import torch.nn as nn
|
| 5 |
+
import torch.nn.functional as F
|
| 6 |
+
|
| 7 |
+
from .parallel import get_parallel_state, exchange_borders
|
| 8 |
+
|
| 9 |
+
|
| 10 |
+
|
| 11 |
+
|
| 12 |
+
class BaseConv3d(nn.Conv3d):
|
| 13 |
+
def __init__(
|
| 14 |
+
self,
|
| 15 |
+
in_channels,
|
| 16 |
+
out_channels,
|
| 17 |
+
kernel_size,
|
| 18 |
+
stride=1,
|
| 19 |
+
padding=0,
|
| 20 |
+
bias=True,
|
| 21 |
+
padding_mode="zeros",
|
| 22 |
+
padding_mode_t=None,
|
| 23 |
+
causal=True,
|
| 24 |
+
):
|
| 25 |
+
super().__init__(
|
| 26 |
+
in_channels,
|
| 27 |
+
out_channels,
|
| 28 |
+
kernel_size=kernel_size,
|
| 29 |
+
stride=stride,
|
| 30 |
+
padding=padding,
|
| 31 |
+
bias=bias,
|
| 32 |
+
padding_mode=padding_mode,
|
| 33 |
+
)
|
| 34 |
+
padding_mode = "constant" if padding_mode == "zeros" else padding_mode
|
| 35 |
+
padding_mode_t = "constant" if padding_mode_t == "zeros" else padding_mode_t
|
| 36 |
+
self.pad_mode = padding_mode
|
| 37 |
+
self.pad_mode_t = padding_mode_t or ("constant" if causal else "replicate")
|
| 38 |
+
self.causal = causal
|
| 39 |
+
|
| 40 |
+
def _apply_temporal_padding(self, x):
|
| 41 |
+
B, C, D, H, W = x.shape
|
| 42 |
+
if D > 1:
|
| 43 |
+
pad_size = (
|
| 44 |
+
0,
|
| 45 |
+
0,
|
| 46 |
+
0,
|
| 47 |
+
0,
|
| 48 |
+
self.padding[0] * 2 if self.causal else self.padding[0],
|
| 49 |
+
0 if self.causal else self.padding[0],
|
| 50 |
+
)
|
| 51 |
+
return F.pad(x, pad_size, mode=self.pad_mode_t)
|
| 52 |
+
else:
|
| 53 |
+
if self.pad_mode_t == "constant":
|
| 54 |
+
assert self.causal, "Zeros padding is only supported for causal mode"
|
| 55 |
+
zeros = torch.zeros_like(x[:, :, :1, :, :]).expand(
|
| 56 |
+
-1, -1, self.kernel_size[0] - 1, -1, -1
|
| 57 |
+
)
|
| 58 |
+
return torch.cat([zeros, x], dim=2)
|
| 59 |
+
else:
|
| 60 |
+
return x.expand(-1, -1, self.kernel_size[0], -1, -1)
|
| 61 |
+
|
| 62 |
+
def _apply_padding(self, x):
|
| 63 |
+
if sum(self.padding) == 0:
|
| 64 |
+
return x
|
| 65 |
+
|
| 66 |
+
x = F.pad(
|
| 67 |
+
x,
|
| 68 |
+
(self.padding[2], self.padding[2], self.padding[1], self.padding[1], 0, 0),
|
| 69 |
+
mode=self.pad_mode,
|
| 70 |
+
)
|
| 71 |
+
|
| 72 |
+
x = self._apply_temporal_padding(x)
|
| 73 |
+
return x
|
| 74 |
+
|
| 75 |
+
def forward(self, x):
|
| 76 |
+
if sum(self.padding) == 0:
|
| 77 |
+
return super().forward(x)
|
| 78 |
+
|
| 79 |
+
x = self._apply_padding(x)
|
| 80 |
+
return F.conv3d(
|
| 81 |
+
x,
|
| 82 |
+
self.weight,
|
| 83 |
+
self.bias,
|
| 84 |
+
stride=self.stride,
|
| 85 |
+
padding=0,
|
| 86 |
+
dilation=self.dilation,
|
| 87 |
+
)
|
| 88 |
+
|
| 89 |
+
|
| 90 |
+
class SpatialParallelConv3d(BaseConv3d):
|
| 91 |
+
def __init__(
|
| 92 |
+
self,
|
| 93 |
+
in_channels,
|
| 94 |
+
out_channels,
|
| 95 |
+
kernel_size,
|
| 96 |
+
stride=1,
|
| 97 |
+
padding=0,
|
| 98 |
+
bias=True,
|
| 99 |
+
padding_mode="zeros",
|
| 100 |
+
padding_mode_t=None,
|
| 101 |
+
causal=True,
|
| 102 |
+
):
|
| 103 |
+
super().__init__(
|
| 104 |
+
in_channels,
|
| 105 |
+
out_channels,
|
| 106 |
+
kernel_size=kernel_size,
|
| 107 |
+
stride=stride,
|
| 108 |
+
padding=padding,
|
| 109 |
+
bias=bias,
|
| 110 |
+
padding_mode=padding_mode,
|
| 111 |
+
padding_mode_t=padding_mode_t,
|
| 112 |
+
causal=causal,
|
| 113 |
+
)
|
| 114 |
+
self.spatial_parallel = False
|
| 115 |
+
self.chunk_dim = -1
|
| 116 |
+
|
| 117 |
+
def _exchange_borders(self, x, sp_rank, sp_size):
|
| 118 |
+
if self.chunk_dim == -1:
|
| 119 |
+
pad = self.padding[2]
|
| 120 |
+
elif self.chunk_dim == -2:
|
| 121 |
+
pad = self.padding[1]
|
| 122 |
+
else:
|
| 123 |
+
raise ValueError(f"Invalid chunk dimension: {self.chunk_dim}")
|
| 124 |
+
|
| 125 |
+
if pad == 0:
|
| 126 |
+
return x
|
| 127 |
+
|
| 128 |
+
local_process_group = get_parallel_state()["sp_process_group"]
|
| 129 |
+
return exchange_borders(
|
| 130 |
+
x,
|
| 131 |
+
pad,
|
| 132 |
+
self.pad_mode,
|
| 133 |
+
sp_rank,
|
| 134 |
+
sp_size,
|
| 135 |
+
local_process_group,
|
| 136 |
+
dim=self.chunk_dim,
|
| 137 |
+
)
|
| 138 |
+
|
| 139 |
+
def _apply_padding(self, x):
|
| 140 |
+
if not self.spatial_parallel:
|
| 141 |
+
return super()._apply_padding(x)
|
| 142 |
+
|
| 143 |
+
state = get_parallel_state()
|
| 144 |
+
|
| 145 |
+
x = self._exchange_borders(x, state["sp_rank"], state["sp_size"])
|
| 146 |
+
|
| 147 |
+
if self.chunk_dim == -1:
|
| 148 |
+
x = F.pad(
|
| 149 |
+
x, (0, 0, self.padding[1], self.padding[1], 0, 0), mode=self.pad_mode
|
| 150 |
+
)
|
| 151 |
+
elif self.chunk_dim == -2:
|
| 152 |
+
x = F.pad(
|
| 153 |
+
x, (self.padding[2], self.padding[2], 0, 0, 0, 0), mode=self.pad_mode
|
| 154 |
+
)
|
| 155 |
+
else:
|
| 156 |
+
raise ValueError(f"Invalid chunk dimension: {self.chunk_dim}")
|
| 157 |
+
|
| 158 |
+
x = self._apply_temporal_padding(x)
|
| 159 |
+
return x
|
FL2VA/video_vae/flash.py
ADDED
|
@@ -0,0 +1,178 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 2 |
+
# Torch-native attention implemented with PyTorch SDPA instead of FA4/CUTLASS.
|
| 3 |
+
import os
|
| 4 |
+
from contextlib import nullcontext
|
| 5 |
+
|
| 6 |
+
import torch
|
| 7 |
+
import torch.nn.functional as F
|
| 8 |
+
|
| 9 |
+
|
| 10 |
+
_BLOCK_CAUSAL_MASK_MOD_CACHE = {}
|
| 11 |
+
|
| 12 |
+
|
| 13 |
+
|
| 14 |
+
|
| 15 |
+
def _as_bool_mask(mask, *, device):
|
| 16 |
+
if not isinstance(mask, torch.Tensor):
|
| 17 |
+
mask = torch.as_tensor(mask, device=device)
|
| 18 |
+
return mask.to(device=device, dtype=torch.bool)
|
| 19 |
+
|
| 20 |
+
|
| 21 |
+
def _ensure_nonempty_rows(mask):
|
| 22 |
+
if mask.numel() == 0 or mask.shape[-1] == 0:
|
| 23 |
+
return mask
|
| 24 |
+
empty = ~mask.any(dim=-1)
|
| 25 |
+
if empty.any():
|
| 26 |
+
mask = mask.clone()
|
| 27 |
+
mask[..., 0] |= empty
|
| 28 |
+
return mask
|
| 29 |
+
|
| 30 |
+
|
| 31 |
+
def _sdpa_kernel_context():
|
| 32 |
+
backend_name = os.environ.get("MINIMAX_H3_TORCH_SDPA_BACKEND", "auto").lower()
|
| 33 |
+
if backend_name in {"", "auto", "default"}:
|
| 34 |
+
return nullcontext()
|
| 35 |
+
|
| 36 |
+
from torch.nn.attention import SDPBackend, sdpa_kernel
|
| 37 |
+
|
| 38 |
+
backends = {
|
| 39 |
+
"math": SDPBackend.MATH,
|
| 40 |
+
"flash": SDPBackend.FLASH_ATTENTION,
|
| 41 |
+
"flash_attention": SDPBackend.FLASH_ATTENTION,
|
| 42 |
+
"efficient": SDPBackend.EFFICIENT_ATTENTION,
|
| 43 |
+
"mem_efficient": SDPBackend.EFFICIENT_ATTENTION,
|
| 44 |
+
"cudnn": SDPBackend.CUDNN_ATTENTION,
|
| 45 |
+
"cudnn_attention": SDPBackend.CUDNN_ATTENTION,
|
| 46 |
+
}
|
| 47 |
+
if backend_name not in backends:
|
| 48 |
+
raise ValueError(
|
| 49 |
+
"MINIMAX_H3_TORCH_SDPA_BACKEND must be one of "
|
| 50 |
+
f"{sorted([*backends, 'auto', 'default'])}, got {backend_name!r}"
|
| 51 |
+
)
|
| 52 |
+
return sdpa_kernel(backends=[backends[backend_name]])
|
| 53 |
+
|
| 54 |
+
|
| 55 |
+
def _sdpa_attention(query, key, value, causal=False, attn_mask=None):
|
| 56 |
+
# query/key/value arrive as [B, S, H, D]; PyTorch SDPA expects
|
| 57 |
+
# [B, H, S, D].
|
| 58 |
+
q = query.transpose(1, 2)
|
| 59 |
+
k = key.transpose(1, 2)
|
| 60 |
+
v = value.transpose(1, 2)
|
| 61 |
+
if attn_mask is not None and attn_mask.dim() == 3:
|
| 62 |
+
attn_mask = attn_mask.unsqueeze(0)
|
| 63 |
+
with _sdpa_kernel_context():
|
| 64 |
+
out = F.scaled_dot_product_attention(
|
| 65 |
+
q,
|
| 66 |
+
k,
|
| 67 |
+
v,
|
| 68 |
+
attn_mask=attn_mask,
|
| 69 |
+
dropout_p=0.0,
|
| 70 |
+
is_causal=causal,
|
| 71 |
+
)
|
| 72 |
+
return out.transpose(1, 2).nan_to_num(0.0)
|
| 73 |
+
|
| 74 |
+
|
| 75 |
+
def _mask_mod_to_dense(mask_mod, batch, heads, q_len, kv_len, device, aux_tensors=None):
|
| 76 |
+
q_idx = torch.arange(q_len, device=device).view(q_len, 1)
|
| 77 |
+
kv_idx = torch.arange(kv_len, device=device).view(1, kv_len)
|
| 78 |
+
dense = torch.empty((batch, heads, q_len, kv_len), dtype=torch.bool, device=device)
|
| 79 |
+
for b in range(batch):
|
| 80 |
+
b_idx = torch.tensor(b, device=device)
|
| 81 |
+
for h in range(heads):
|
| 82 |
+
h_idx = torch.tensor(h, device=device)
|
| 83 |
+
mask = mask_mod(b_idx, h_idx, q_idx, kv_idx, None, aux_tensors)
|
| 84 |
+
dense[b, h] = _as_bool_mask(mask, device=device)
|
| 85 |
+
return _ensure_nonempty_rows(dense)
|
| 86 |
+
|
| 87 |
+
|
| 88 |
+
#########################################################
|
| 89 |
+
# Block causal attention
|
| 90 |
+
#########################################################
|
| 91 |
+
|
| 92 |
+
|
| 93 |
+
def make_block_causal_mask_mod(num_tokens, block_size, num_special=0, suffix=False):
|
| 94 |
+
if num_tokens < 0:
|
| 95 |
+
raise ValueError(f"num_tokens must be non-negative, got {num_tokens}")
|
| 96 |
+
if block_size <= 0:
|
| 97 |
+
raise ValueError(f"block_size must be positive, got {block_size}")
|
| 98 |
+
if num_special < 0:
|
| 99 |
+
raise ValueError(f"num_special must be non-negative, got {num_special}")
|
| 100 |
+
|
| 101 |
+
cache_key = (num_tokens, block_size, num_special, suffix)
|
| 102 |
+
if cache_key in _BLOCK_CAUSAL_MASK_MOD_CACHE:
|
| 103 |
+
return _BLOCK_CAUSAL_MASK_MOD_CACHE[cache_key]
|
| 104 |
+
|
| 105 |
+
if suffix:
|
| 106 |
+
|
| 107 |
+
def mask_mod(b, h, q_idx, kv_idx, seqlen_info, aux_tensors):
|
| 108 |
+
del b, h, seqlen_info, aux_tensors
|
| 109 |
+
q_is_special = q_idx >= num_tokens
|
| 110 |
+
kv_is_special = kv_idx >= num_tokens
|
| 111 |
+
return q_is_special | kv_is_special | (
|
| 112 |
+
q_idx // block_size >= kv_idx // block_size
|
| 113 |
+
)
|
| 114 |
+
|
| 115 |
+
else:
|
| 116 |
+
|
| 117 |
+
def mask_mod(b, h, q_idx, kv_idx, seqlen_info, aux_tensors):
|
| 118 |
+
del b, h, seqlen_info, aux_tensors
|
| 119 |
+
q_is_special = q_idx < num_special
|
| 120 |
+
kv_is_special = kv_idx < num_special
|
| 121 |
+
q_block_idx = (q_idx - num_special) // block_size
|
| 122 |
+
kv_block_idx = (kv_idx - num_special) // block_size
|
| 123 |
+
return q_is_special | kv_is_special | (q_block_idx >= kv_block_idx)
|
| 124 |
+
|
| 125 |
+
mask_mod.block_sparse_cache_key = (
|
| 126 |
+
"block_causal",
|
| 127 |
+
num_tokens,
|
| 128 |
+
block_size,
|
| 129 |
+
num_special,
|
| 130 |
+
suffix,
|
| 131 |
+
)
|
| 132 |
+
_BLOCK_CAUSAL_MASK_MOD_CACHE[cache_key] = mask_mod
|
| 133 |
+
return mask_mod
|
| 134 |
+
|
| 135 |
+
|
| 136 |
+
|
| 137 |
+
|
| 138 |
+
|
| 139 |
+
|
| 140 |
+
#########################################################
|
| 141 |
+
# Public entry point
|
| 142 |
+
#########################################################
|
| 143 |
+
|
| 144 |
+
|
| 145 |
+
@torch.compiler.disable
|
| 146 |
+
def flash_attn(
|
| 147 |
+
query: torch.Tensor,
|
| 148 |
+
key: torch.Tensor,
|
| 149 |
+
value: torch.Tensor,
|
| 150 |
+
causal: bool = False,
|
| 151 |
+
mask_mod=None,
|
| 152 |
+
block_sparse=None,
|
| 153 |
+
aux_tensors=None,
|
| 154 |
+
) -> torch.Tensor:
|
| 155 |
+
use_masked = mask_mod is not None or block_sparse is not None
|
| 156 |
+
|
| 157 |
+
if block_sparse is not None and mask_mod is None:
|
| 158 |
+
raise ValueError("block_sparse requires mask_mod")
|
| 159 |
+
if causal and mask_mod is not None:
|
| 160 |
+
raise ValueError("causal must be encoded in mask_mod when using masked attention")
|
| 161 |
+
if aux_tensors is not None and not use_masked:
|
| 162 |
+
raise ValueError("aux_tensors is only supported with masked attention")
|
| 163 |
+
|
| 164 |
+
if use_masked:
|
| 165 |
+
batch, q_len, heads, _ = query.shape
|
| 166 |
+
kv_len = key.shape[1]
|
| 167 |
+
dense_mask = _mask_mod_to_dense(
|
| 168 |
+
mask_mod,
|
| 169 |
+
batch,
|
| 170 |
+
heads,
|
| 171 |
+
q_len,
|
| 172 |
+
kv_len,
|
| 173 |
+
query.device,
|
| 174 |
+
aux_tensors=aux_tensors,
|
| 175 |
+
)
|
| 176 |
+
return _sdpa_attention(query, key, value, attn_mask=dense_mask)
|
| 177 |
+
|
| 178 |
+
return _sdpa_attention(query, key, value, causal=causal)
|
FL2VA/video_vae/func.py
ADDED
|
@@ -0,0 +1,163 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 2 |
+
# Token-id and rotary-embedding helpers for the MiniMax H3 visual VAE.
|
| 3 |
+
import os
|
| 4 |
+
import torch
|
| 5 |
+
from typing import Tuple
|
| 6 |
+
|
| 7 |
+
from diffusers.utils import logging
|
| 8 |
+
|
| 9 |
+
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
| 10 |
+
|
| 11 |
+
|
| 12 |
+
def create_token_ids(patch_dims, device, dtype, id_type="length_normalized", flatten=True):
|
| 13 |
+
coords_list = []
|
| 14 |
+
|
| 15 |
+
if isinstance(id_type, str):
|
| 16 |
+
id_type_list = [id_type] * len(patch_dims)
|
| 17 |
+
elif isinstance(id_type, list):
|
| 18 |
+
id_type_list = id_type
|
| 19 |
+
if len(id_type_list) != len(patch_dims):
|
| 20 |
+
raise ValueError("id_type list must match patch_dims")
|
| 21 |
+
else:
|
| 22 |
+
raise ValueError("id_type must be a string or a list")
|
| 23 |
+
|
| 24 |
+
if "area_normalized" in id_type_list or id_type == "area_normalized":
|
| 25 |
+
raise NotImplementedError(
|
| 26 |
+
"area_normalized id_type is not supported in this inference-only bundle"
|
| 27 |
+
)
|
| 28 |
+
|
| 29 |
+
for _dim_size, _id_type in zip(patch_dims, id_type_list):
|
| 30 |
+
if isinstance(_dim_size, torch.Tensor):
|
| 31 |
+
coords_list.append(_dim_size.to(device=device, dtype=dtype))
|
| 32 |
+
continue
|
| 33 |
+
|
| 34 |
+
if _id_type == "length_normalized":
|
| 35 |
+
coords = torch.arange(0.5, _dim_size, dtype=dtype, device=device)
|
| 36 |
+
coords = coords / _dim_size
|
| 37 |
+
coords = 2.0 * coords - 1.0
|
| 38 |
+
else:
|
| 39 |
+
coords = torch.arange(_dim_size, dtype=dtype, device=device)
|
| 40 |
+
|
| 41 |
+
coords_list.append(coords)
|
| 42 |
+
|
| 43 |
+
coords = torch.stack(torch.meshgrid(*coords_list, indexing="ij"), dim=-1)
|
| 44 |
+
if flatten:
|
| 45 |
+
coords = coords.flatten(0, len(patch_dims) - 1)
|
| 46 |
+
|
| 47 |
+
return coords.unsqueeze(0)
|
| 48 |
+
|
| 49 |
+
|
| 50 |
+
def _env_flag(name, default="0"):
|
| 51 |
+
value = os.environ.get(name, default)
|
| 52 |
+
return str(value).strip().lower() in ("1", "true", "yes", "on")
|
| 53 |
+
|
| 54 |
+
|
| 55 |
+
def _env_optional_bool(name, default=""):
|
| 56 |
+
value = str(os.environ.get(name, default)).strip().lower()
|
| 57 |
+
if value in ("", "default", "auto", "none", "unset"):
|
| 58 |
+
return None
|
| 59 |
+
return value not in ("0", "false", "no", "off", "disabled")
|
| 60 |
+
|
| 61 |
+
|
| 62 |
+
def _vit_torch_compile_kwargs(prefix):
|
| 63 |
+
kwargs = {}
|
| 64 |
+
backend = os.environ.get(f"{prefix}_BACKEND", "inductor").strip()
|
| 65 |
+
mode = os.environ.get(f"{prefix}_MODE", "reduce-overhead").strip()
|
| 66 |
+
if backend and backend.lower() not in ("default", "none"):
|
| 67 |
+
kwargs["backend"] = backend
|
| 68 |
+
if mode and mode.lower() not in ("default", "none"):
|
| 69 |
+
kwargs["mode"] = mode
|
| 70 |
+
kwargs["fullgraph"] = _env_flag(f"{prefix}_FULLGRAPH", "0")
|
| 71 |
+
dynamic = _env_optional_bool(f"{prefix}_DYNAMIC")
|
| 72 |
+
if dynamic is not None:
|
| 73 |
+
kwargs["dynamic"] = dynamic
|
| 74 |
+
return kwargs
|
| 75 |
+
|
| 76 |
+
|
| 77 |
+
def _rotate_half(x: torch.Tensor) -> torch.Tensor:
|
| 78 |
+
x1, x2 = torch.chunk(x, 2, dim=-1)
|
| 79 |
+
return torch.cat((-x2, x1), dim=-1)
|
| 80 |
+
|
| 81 |
+
|
| 82 |
+
def _apply_rotary_pos_emb_impl(
|
| 83 |
+
t: torch.Tensor, rotary_pos_emb: Tuple[torch.Tensor, torch.Tensor]
|
| 84 |
+
) -> torch.Tensor:
|
| 85 |
+
cos, sin = rotary_pos_emb
|
| 86 |
+
|
| 87 |
+
if cos.dim() != 4:
|
| 88 |
+
raise ValueError(f"cos must be [B, N, 1, D], got {cos.shape}")
|
| 89 |
+
|
| 90 |
+
cos = cos.to(t.dtype)
|
| 91 |
+
sin = sin.to(t.dtype)
|
| 92 |
+
|
| 93 |
+
rot_dim = cos.shape[-1]
|
| 94 |
+
t_dim = t.shape[-1]
|
| 95 |
+
|
| 96 |
+
if rot_dim < t_dim:
|
| 97 |
+
t_rot, t_pass = t[..., :rot_dim], t[..., rot_dim:]
|
| 98 |
+
t_rot = (t_rot * cos) + (_rotate_half(t_rot) * sin)
|
| 99 |
+
t = torch.cat((t_rot, t_pass), dim=-1)
|
| 100 |
+
else:
|
| 101 |
+
t = (t * cos) + (_rotate_half(t) * sin)
|
| 102 |
+
|
| 103 |
+
return t
|
| 104 |
+
|
| 105 |
+
_COMPILED_APPLY_ROTARY_POS_EMB = None
|
| 106 |
+
_APPLY_ROTARY_POS_EMB_COMPILE_DISABLED = False
|
| 107 |
+
|
| 108 |
+
|
| 109 |
+
def _get_apply_rotary_pos_emb_impl():
|
| 110 |
+
global _COMPILED_APPLY_ROTARY_POS_EMB, _APPLY_ROTARY_POS_EMB_COMPILE_DISABLED
|
| 111 |
+
if _APPLY_ROTARY_POS_EMB_COMPILE_DISABLED or not _env_flag(
|
| 112 |
+
"MINIMAX_H3_VAE_DECODER_VIT_ROPE_TORCH_COMPILE", "0"
|
| 113 |
+
):
|
| 114 |
+
return _apply_rotary_pos_emb_impl
|
| 115 |
+
if _COMPILED_APPLY_ROTARY_POS_EMB is not None:
|
| 116 |
+
return _COMPILED_APPLY_ROTARY_POS_EMB
|
| 117 |
+
if not hasattr(torch, "compile"):
|
| 118 |
+
message = "torch.compile is unavailable; falling back to eager ViT rotary embedding"
|
| 119 |
+
if _env_flag("MINIMAX_H3_VAE_DECODER_VIT_ROPE_TORCH_COMPILE_FATAL", "0"):
|
| 120 |
+
raise RuntimeError(message)
|
| 121 |
+
logger.warning(f"[ViTRope] {message}")
|
| 122 |
+
_APPLY_ROTARY_POS_EMB_COMPILE_DISABLED = True
|
| 123 |
+
return _apply_rotary_pos_emb_impl
|
| 124 |
+
|
| 125 |
+
kwargs = _vit_torch_compile_kwargs("MINIMAX_H3_VAE_DECODER_VIT_ROPE_TORCH_COMPILE")
|
| 126 |
+
try:
|
| 127 |
+
_COMPILED_APPLY_ROTARY_POS_EMB = torch.compile(
|
| 128 |
+
_apply_rotary_pos_emb_impl, **kwargs
|
| 129 |
+
)
|
| 130 |
+
logger.info(f"[ViTRope] torch.compile enabled kwargs={kwargs}")
|
| 131 |
+
except Exception as exc:
|
| 132 |
+
if _env_flag("MINIMAX_H3_VAE_DECODER_VIT_ROPE_TORCH_COMPILE_FATAL", "0"):
|
| 133 |
+
raise
|
| 134 |
+
logger.warning(
|
| 135 |
+
f"[ViTRope] torch.compile setup failed: {type(exc).__name__}: {exc}; "
|
| 136 |
+
"falling back to eager"
|
| 137 |
+
)
|
| 138 |
+
_APPLY_ROTARY_POS_EMB_COMPILE_DISABLED = True
|
| 139 |
+
_COMPILED_APPLY_ROTARY_POS_EMB = None
|
| 140 |
+
return _apply_rotary_pos_emb_impl
|
| 141 |
+
return _COMPILED_APPLY_ROTARY_POS_EMB
|
| 142 |
+
|
| 143 |
+
|
| 144 |
+
def apply_rotary_pos_emb(
|
| 145 |
+
t: torch.Tensor, rotary_pos_emb: Tuple[torch.Tensor, torch.Tensor]
|
| 146 |
+
) -> torch.Tensor:
|
| 147 |
+
global _COMPILED_APPLY_ROTARY_POS_EMB, _APPLY_ROTARY_POS_EMB_COMPILE_DISABLED
|
| 148 |
+
fn = _get_apply_rotary_pos_emb_impl()
|
| 149 |
+
try:
|
| 150 |
+
return fn(t, rotary_pos_emb)
|
| 151 |
+
except Exception as exc:
|
| 152 |
+
if (
|
| 153 |
+
fn is _COMPILED_APPLY_ROTARY_POS_EMB
|
| 154 |
+
and not _env_flag("MINIMAX_H3_VAE_DECODER_VIT_ROPE_TORCH_COMPILE_FATAL", "0")
|
| 155 |
+
):
|
| 156 |
+
logger.warning(
|
| 157 |
+
f"[ViTRope] compiled call failed: {type(exc).__name__}: {exc}; "
|
| 158 |
+
"disabling compile and retrying eager"
|
| 159 |
+
)
|
| 160 |
+
_APPLY_ROTARY_POS_EMB_COMPILE_DISABLED = True
|
| 161 |
+
_COMPILED_APPLY_ROTARY_POS_EMB = None
|
| 162 |
+
return _apply_rotary_pos_emb_impl(t, rotary_pos_emb)
|
| 163 |
+
raise
|
FL2VA/video_vae/klvae.py
ADDED
|
@@ -0,0 +1,1258 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 2 |
+
# MiniMax H3 visual VAE: 3D causal CNN encoder + ViT3D decoder (inference-only bundle).
|
| 3 |
+
import os
|
| 4 |
+
import math
|
| 5 |
+
import numpy as np
|
| 6 |
+
import torch
|
| 7 |
+
import torch.nn as nn
|
| 8 |
+
import torch.distributed as dist
|
| 9 |
+
from typing import List, Union
|
| 10 |
+
from PIL import Image
|
| 11 |
+
from contextlib import nullcontext
|
| 12 |
+
from diffusers.models import ModelMixin
|
| 13 |
+
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
| 14 |
+
from diffusers.loaders.single_file_model import FromOriginalModelMixin
|
| 15 |
+
from diffusers.utils import logging
|
| 16 |
+
|
| 17 |
+
from .parallel import get_parallel_state, all_gather_var_shape
|
| 18 |
+
from .utils import apply_spatial_parallel
|
| 19 |
+
from .normalize import get_normalize_transform, get_denormalize_transform
|
| 20 |
+
from .vae_vit import ViT3DDecoder
|
| 21 |
+
from .vae_cnn import EncoderFCN3D
|
| 22 |
+
from .vae_module import DiagonalGaussianDistribution, ClsTokenAggregator
|
| 23 |
+
from .vae_processor import VAEProcessor
|
| 24 |
+
|
| 25 |
+
|
| 26 |
+
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
| 27 |
+
|
| 28 |
+
|
| 29 |
+
def _resolve_temporal_cat_dtype():
|
| 30 |
+
raw = os.environ.get("MINIMAX_H3_VAE_DECODER_TEMPORAL_CAT_DTYPE", "").strip().lower()
|
| 31 |
+
if raw in ("", "0", "false", "no", "off", "none", "keep", "default"):
|
| 32 |
+
return None
|
| 33 |
+
mapping = {
|
| 34 |
+
"fp16": torch.float16,
|
| 35 |
+
"float16": torch.float16,
|
| 36 |
+
"half": torch.float16,
|
| 37 |
+
"bf16": torch.bfloat16,
|
| 38 |
+
"bfloat16": torch.bfloat16,
|
| 39 |
+
"fp32": torch.float32,
|
| 40 |
+
"float32": torch.float32,
|
| 41 |
+
}
|
| 42 |
+
if raw not in mapping:
|
| 43 |
+
raise ValueError(
|
| 44 |
+
"MINIMAX_H3_VAE_DECODER_TEMPORAL_CAT_DTYPE must be one of "
|
| 45 |
+
"fp16|bf16|fp32|keep, got %r" % raw
|
| 46 |
+
)
|
| 47 |
+
return mapping[raw]
|
| 48 |
+
|
| 49 |
+
|
| 50 |
+
def _resolve_temporal_stream_cat():
|
| 51 |
+
raw = os.environ.get("MINIMAX_H3_VAE_DECODER_STREAM_TEMPORAL_CAT", "1").strip().lower()
|
| 52 |
+
return raw not in ("0", "false", "no", "off", "disable", "disabled")
|
| 53 |
+
|
| 54 |
+
|
| 55 |
+
class AutoencoderKL(ModelMixin, ConfigMixin, FromOriginalModelMixin):
|
| 56 |
+
r"""
|
| 57 |
+
Abstract shared base for the MiniMax H3 visual VAE.
|
| 58 |
+
|
| 59 |
+
This class only carries the shared inference machinery (temporal
|
| 60 |
+
chunking, tiling, encode/decode entry points). Instantiate the concrete
|
| 61 |
+
subclass ``AutoencoderKLLegacy`` via ``from_pretrained`` instead.
|
| 62 |
+
"""
|
| 63 |
+
|
| 64 |
+
_supports_gradient_checkpointing = True
|
| 65 |
+
_compilable_modules = ["encoder", "decoder"]
|
| 66 |
+
_deprecated_kwargs = [
|
| 67 |
+
"clip_length",
|
| 68 |
+
"token_drop",
|
| 69 |
+
"isolated_first_frame",
|
| 70 |
+
"isolated_last_frame",
|
| 71 |
+
"isolated_key_frame",
|
| 72 |
+
"encoder_tiling",
|
| 73 |
+
"decoder_tiling",
|
| 74 |
+
"parallel_tiling",
|
| 75 |
+
"stack_tiling",
|
| 76 |
+
"tile_size",
|
| 77 |
+
"tile_overlap_min",
|
| 78 |
+
"decoder_tile_size",
|
| 79 |
+
"decoder_tile_overlap_min",
|
| 80 |
+
"latent_patch_size",
|
| 81 |
+
"crop_mode",
|
| 82 |
+
"encoder_parallel",
|
| 83 |
+
"decoder_parallel",
|
| 84 |
+
"chunk_dim",
|
| 85 |
+
] # legacy config keys accepted by from_pretrained for checkpoint compatibility
|
| 86 |
+
|
| 87 |
+
|
| 88 |
+
def _set_gradient_checkpointing(self, module, value=False):
|
| 89 |
+
if hasattr(module, "gradient_checkpointing"):
|
| 90 |
+
module.gradient_checkpointing = value
|
| 91 |
+
|
| 92 |
+
def _freeze_nested_module(self, module_path):
|
| 93 |
+
parts = module_path.split(".")
|
| 94 |
+
module = self
|
| 95 |
+
for part in parts:
|
| 96 |
+
module = getattr(module, part)
|
| 97 |
+
module.requires_grad_(False)
|
| 98 |
+
|
| 99 |
+
def setup_forward(self, **kwargs):
|
| 100 |
+
self.clip_length = kwargs.get("clip_length", 17)
|
| 101 |
+
self.token_drop = kwargs.get("token_drop", 0)
|
| 102 |
+
self.frame_drop = self.token_drop * self.vae_ratio_t
|
| 103 |
+
self.frame_pre_padding = (-self.clip_length) % self.vae_ratio_t
|
| 104 |
+
self.tokens_chunk_size = math.ceil(self.clip_length / self.vae_ratio_t)
|
| 105 |
+
self.token_overlap = (-self.token_drop) % self.tokens_chunk_size
|
| 106 |
+
self.frame_overlap = max(self.token_overlap * self.vae_ratio_t - self.frame_pre_padding, 0)
|
| 107 |
+
self.isolated_first_frame = kwargs.get("isolated_first_frame", False)
|
| 108 |
+
self.isolated_last_frame = kwargs.get("isolated_last_frame", False)
|
| 109 |
+
self.isolated_key_frame = kwargs.get("isolated_key_frame", False)
|
| 110 |
+
|
| 111 |
+
self.encoder_tiling = kwargs.get("encoder_tiling", False)
|
| 112 |
+
self.decoder_tiling = kwargs.get("decoder_tiling", False)
|
| 113 |
+
self.stack_tiling = kwargs.get("stack_tiling", False)
|
| 114 |
+
self.tile_size = kwargs.get("tile_size", 256)
|
| 115 |
+
self.tile_overlap_min = kwargs.get("tile_overlap_min", 64)
|
| 116 |
+
self.decoder_tile_size = kwargs.get("decoder_tile_size", self.tile_size)
|
| 117 |
+
self.decoder_tile_overlap_min = kwargs.get("decoder_tile_overlap_min", self.tile_overlap_min)
|
| 118 |
+
self.latent_patch_size = kwargs.get("latent_patch_size", 1)
|
| 119 |
+
self.crop_mode = kwargs.get("crop_mode", "top_left")
|
| 120 |
+
self.pixel_norm_type = kwargs.get("pixel_norm_type", "imagenet")
|
| 121 |
+
|
| 122 |
+
# spatial parallel mode
|
| 123 |
+
if hasattr(self, "_sp_initialized"):
|
| 124 |
+
if (
|
| 125 |
+
kwargs.get("chunk_dim", -1) != self.chunk_dim
|
| 126 |
+
or kwargs.get("encoder_parallel", False) != self.encoder_parallel
|
| 127 |
+
or kwargs.get("decoder_parallel", False) != self.decoder_parallel
|
| 128 |
+
or kwargs.get("parallel_tiling", False) != self.parallel_tiling
|
| 129 |
+
):
|
| 130 |
+
logger.warning(
|
| 131 |
+
"Do not support changing parallel schema after initialization"
|
| 132 |
+
)
|
| 133 |
+
else:
|
| 134 |
+
self.chunk_dim = kwargs.get("chunk_dim", -1)
|
| 135 |
+
self.encoder_parallel = kwargs.get("encoder_parallel", False)
|
| 136 |
+
self.decoder_parallel = kwargs.get("decoder_parallel", False)
|
| 137 |
+
self.parallel_tiling = kwargs.get("parallel_tiling", False)
|
| 138 |
+
self._sp_initialized = True
|
| 139 |
+
|
| 140 |
+
processor_kwargs = {
|
| 141 |
+
"vae_ratio": self.vae_ratio,
|
| 142 |
+
"vae_ratio_t": self.vae_ratio_t,
|
| 143 |
+
"clip_length": self.clip_length,
|
| 144 |
+
"frame_overlap": self.frame_overlap,
|
| 145 |
+
"token_overlap": self.token_overlap,
|
| 146 |
+
"tokens_chunk_size": self.tokens_chunk_size,
|
| 147 |
+
"isolated_last_frame": self.isolated_last_frame,
|
| 148 |
+
"latent_patch_size": self.latent_patch_size,
|
| 149 |
+
"crop_mode": self.crop_mode,
|
| 150 |
+
"pixel_norm_type": self.pixel_norm_type,
|
| 151 |
+
"transform": self.transform,
|
| 152 |
+
"transform_rev": self.transform_rev,
|
| 153 |
+
"use_3d_conv": self.use_3d_conv,
|
| 154 |
+
}
|
| 155 |
+
if hasattr(self, "processor"):
|
| 156 |
+
for key, value in processor_kwargs.items():
|
| 157 |
+
setattr(self.processor, key, value)
|
| 158 |
+
else:
|
| 159 |
+
self.processor = VAEProcessor(**processor_kwargs)
|
| 160 |
+
|
| 161 |
+
def perform_input_slice(self, x, chunk_size_stride=1):
|
| 162 |
+
state = get_parallel_state()
|
| 163 |
+
sp_rank = state["sp_rank"]
|
| 164 |
+
sp_size = state["sp_size"]
|
| 165 |
+
|
| 166 |
+
total_size = x.shape[self.chunk_dim]
|
| 167 |
+
units = total_size // chunk_size_stride
|
| 168 |
+
base_units = units // sp_size
|
| 169 |
+
remainder_units = units % sp_size
|
| 170 |
+
if sp_rank < remainder_units:
|
| 171 |
+
start_units = sp_rank * (base_units + 1)
|
| 172 |
+
end_units = start_units + base_units + 1
|
| 173 |
+
else:
|
| 174 |
+
start_units = sp_rank * base_units + remainder_units
|
| 175 |
+
end_units = start_units + base_units
|
| 176 |
+
start = start_units * chunk_size_stride
|
| 177 |
+
end = end_units * chunk_size_stride
|
| 178 |
+
|
| 179 |
+
slice_indices = [slice(None)] * x.ndim
|
| 180 |
+
slice_indices[self.chunk_dim] = slice(start, end)
|
| 181 |
+
x = x[tuple(slice_indices)].contiguous()
|
| 182 |
+
return x
|
| 183 |
+
|
| 184 |
+
def perform_output_concat(self, x):
|
| 185 |
+
sp_process_group = get_parallel_state()["sp_process_group"]
|
| 186 |
+
gathered = all_gather_var_shape(x, group=sp_process_group)
|
| 187 |
+
x = torch.cat(gathered, dim=self.chunk_dim)
|
| 188 |
+
return x
|
| 189 |
+
|
| 190 |
+
|
| 191 |
+
|
| 192 |
+
def split_tiles(self, input_len, is_decoder=False):
|
| 193 |
+
tile_size = self.decoder_tile_size if is_decoder else self.tile_size
|
| 194 |
+
tile_overlap_min = self.decoder_tile_overlap_min if is_decoder else self.tile_overlap_min
|
| 195 |
+
|
| 196 |
+
if tile_size >= input_len:
|
| 197 |
+
return [0], [input_len], []
|
| 198 |
+
|
| 199 |
+
N = math.ceil(input_len / tile_size)
|
| 200 |
+
while True:
|
| 201 |
+
overlaps = [tile_overlap_min] * (N - 1)
|
| 202 |
+
remaining = tile_size * N - sum(overlaps) - input_len
|
| 203 |
+
|
| 204 |
+
if remaining < 0:
|
| 205 |
+
N += 1
|
| 206 |
+
else:
|
| 207 |
+
break
|
| 208 |
+
|
| 209 |
+
remaining_units = remaining // self.vae_ratio
|
| 210 |
+
for i in range(remaining_units):
|
| 211 |
+
overlaps[i % (N - 1)] += self.vae_ratio
|
| 212 |
+
|
| 213 |
+
tile_start_idx = [0]
|
| 214 |
+
for i in range(N - 1):
|
| 215 |
+
tile_start_idx.append(tile_start_idx[-1] + tile_size - overlaps[i])
|
| 216 |
+
|
| 217 |
+
tile_len = [tile_size] * N
|
| 218 |
+
return tile_start_idx, tile_len, overlaps
|
| 219 |
+
|
| 220 |
+
def blend(
|
| 221 |
+
self, a: torch.Tensor, b: torch.Tensor, blend_extent: int, dim: int
|
| 222 |
+
) -> torch.Tensor:
|
| 223 |
+
blend_extent = min(a.shape[dim], b.shape[dim], blend_extent)
|
| 224 |
+
|
| 225 |
+
positions = torch.arange(blend_extent, device=b.device, dtype=b.dtype)
|
| 226 |
+
weight_a = 1 - positions / blend_extent
|
| 227 |
+
weight_b = positions / blend_extent
|
| 228 |
+
|
| 229 |
+
shape = [1] * a.ndim
|
| 230 |
+
shape[dim] = blend_extent
|
| 231 |
+
weight_a = weight_a.view(shape)
|
| 232 |
+
weight_b = weight_b.view(shape)
|
| 233 |
+
|
| 234 |
+
slice_a = [slice(None)] * a.ndim
|
| 235 |
+
slice_a[dim] = slice(-blend_extent, None)
|
| 236 |
+
a_overlap = a[tuple(slice_a)]
|
| 237 |
+
|
| 238 |
+
slice_b = [slice(None)] * b.ndim
|
| 239 |
+
slice_b[dim] = slice(0, blend_extent)
|
| 240 |
+
b_overlap = b[tuple(slice_b)]
|
| 241 |
+
|
| 242 |
+
blended = a_overlap * weight_a + b_overlap * weight_b
|
| 243 |
+
|
| 244 |
+
if blend_extent < b.shape[dim]:
|
| 245 |
+
slice_b_rest = [slice(None)] * b.ndim
|
| 246 |
+
slice_b_rest[dim] = slice(blend_extent, None)
|
| 247 |
+
b_rest = b[tuple(slice_b_rest)]
|
| 248 |
+
return torch.cat([blended, b_rest], dim=dim)
|
| 249 |
+
else:
|
| 250 |
+
return blended
|
| 251 |
+
|
| 252 |
+
def _all_gather_tiled_results(self, tasks, num_tiles):
|
| 253 |
+
state = get_parallel_state()
|
| 254 |
+
group = state["sp_process_group"]
|
| 255 |
+
sp_size = state["sp_size"]
|
| 256 |
+
sp_rank = state["sp_rank"]
|
| 257 |
+
|
| 258 |
+
if not tasks:
|
| 259 |
+
raise ValueError(f"Found empty tasks on sp rank {sp_rank}")
|
| 260 |
+
|
| 261 |
+
stacked = torch.stack(tasks, dim=0)
|
| 262 |
+
gathered = all_gather_var_shape(stacked, group=group)
|
| 263 |
+
|
| 264 |
+
results = [None] * num_tiles
|
| 265 |
+
for rank, rank_tensors in enumerate(gathered):
|
| 266 |
+
num_rank_tasks = rank_tensors.shape[0]
|
| 267 |
+
for k in range(num_rank_tasks):
|
| 268 |
+
global_idx = k * sp_size + rank
|
| 269 |
+
if global_idx >= num_tiles:
|
| 270 |
+
break
|
| 271 |
+
results[global_idx] = rank_tensors[k]
|
| 272 |
+
|
| 273 |
+
return results
|
| 274 |
+
|
| 275 |
+
def _local_tile_indices(self, num_tiles, sp_rank, sp_size):
|
| 276 |
+
return list(range(sp_rank, num_tiles, sp_size))
|
| 277 |
+
|
| 278 |
+
def _run_tile_tasks(self, tiles, tile_indices, forward_fn, stack_tiling, cls_agg=None):
|
| 279 |
+
if stack_tiling and tile_indices:
|
| 280 |
+
sample_batch_size = tiles[0].shape[0]
|
| 281 |
+
tile_batch = torch.cat([tiles[idx] for idx in tile_indices], dim=0)
|
| 282 |
+
output_batch = forward_fn(tile_batch)
|
| 283 |
+
output_tiles = output_batch.unflatten(
|
| 284 |
+
0, (len(tile_indices), sample_batch_size)
|
| 285 |
+
).unbind(dim=0)
|
| 286 |
+
if cls_agg is not None:
|
| 287 |
+
cls_agg.collect_stacked(len(tile_indices), sample_batch_size)
|
| 288 |
+
return list(output_tiles)
|
| 289 |
+
|
| 290 |
+
tasks = []
|
| 291 |
+
for idx in tile_indices:
|
| 292 |
+
tasks.append(forward_fn(tiles[idx]))
|
| 293 |
+
if cls_agg is not None:
|
| 294 |
+
cls_agg.collect()
|
| 295 |
+
return tasks
|
| 296 |
+
|
| 297 |
+
def tiled_encode(self, x):
|
| 298 |
+
if self.parallel_tiling: # Fast online encoding for large videos
|
| 299 |
+
state = get_parallel_state()
|
| 300 |
+
sp_rank = state["sp_rank"]
|
| 301 |
+
sp_size = state["sp_size"]
|
| 302 |
+
else:
|
| 303 |
+
sp_rank, sp_size = 0, 1
|
| 304 |
+
|
| 305 |
+
height, width = x.shape[-2], x.shape[-1]
|
| 306 |
+
y_idx, y_len, y_overlap = self.split_tiles(height, False)
|
| 307 |
+
x_idx, x_len, x_overlap = self.split_tiles(width, False)
|
| 308 |
+
|
| 309 |
+
i_max, j_max = len(y_idx), len(x_idx)
|
| 310 |
+
num_tiles = i_max * j_max
|
| 311 |
+
|
| 312 |
+
x_tiles = []
|
| 313 |
+
for i, (i_pos, i_len) in enumerate(zip(y_idx, y_len)):
|
| 314 |
+
for j, (j_pos, j_len) in enumerate(zip(x_idx, x_len)):
|
| 315 |
+
tile = x[..., i_pos : i_pos + i_len, j_pos : j_pos + j_len]
|
| 316 |
+
x_tiles.append(tile)
|
| 317 |
+
|
| 318 |
+
with ClsTokenAggregator(self) as agg:
|
| 319 |
+
local_tile_indices = self._local_tile_indices(num_tiles, sp_rank, sp_size)
|
| 320 |
+
stack_tiling = self.stack_tiling and not (
|
| 321 |
+
self.training and getattr(self.encoder, "mask_enabled", False)
|
| 322 |
+
)
|
| 323 |
+
encoded_tasks = self._run_tile_tasks(
|
| 324 |
+
x_tiles, local_tile_indices, self.encode, stack_tiling, agg
|
| 325 |
+
)
|
| 326 |
+
|
| 327 |
+
if sp_size > 1:
|
| 328 |
+
dist.barrier(group=get_parallel_state()["sp_process_group"])
|
| 329 |
+
all_encoded = self._all_gather_tiled_results(encoded_tasks, num_tiles)
|
| 330 |
+
if agg.cls_tokens:
|
| 331 |
+
agg.cls_tokens = self._all_gather_tiled_results(agg.cls_tokens, num_tiles)
|
| 332 |
+
else:
|
| 333 |
+
all_encoded = encoded_tasks
|
| 334 |
+
|
| 335 |
+
rows = [[None for _ in range(j_max)] for _ in range(i_max)]
|
| 336 |
+
for idx, encoded in enumerate(all_encoded):
|
| 337 |
+
i, j = idx // j_max, idx % j_max
|
| 338 |
+
rows[i][j] = encoded.to(x.device)
|
| 339 |
+
|
| 340 |
+
latent_y_overlap = [
|
| 341 |
+
tile_overlap // self.vae_ratio for tile_overlap in y_overlap
|
| 342 |
+
]
|
| 343 |
+
latent_x_overlap = [
|
| 344 |
+
tile_overlap // self.vae_ratio for tile_overlap in x_overlap
|
| 345 |
+
]
|
| 346 |
+
|
| 347 |
+
result_rows = []
|
| 348 |
+
for i, row in enumerate(rows):
|
| 349 |
+
result_row = []
|
| 350 |
+
for j, tile in enumerate(row):
|
| 351 |
+
if i > 0:
|
| 352 |
+
tile = self.blend(rows[i - 1][j], tile, latent_y_overlap[i - 1], dim=-2)
|
| 353 |
+
if j > 0:
|
| 354 |
+
tile = self.blend(row[j - 1], tile, latent_x_overlap[j - 1], dim=-1)
|
| 355 |
+
if i < len(rows) - 1:
|
| 356 |
+
tile = tile[..., : -latent_y_overlap[i], :]
|
| 357 |
+
if j < len(row) - 1:
|
| 358 |
+
tile = tile[..., :, : -latent_x_overlap[j]]
|
| 359 |
+
result_row.append(tile)
|
| 360 |
+
result_rows.append(torch.cat(result_row, dim=-1))
|
| 361 |
+
z = torch.cat(result_rows, dim=-2)
|
| 362 |
+
|
| 363 |
+
return z
|
| 364 |
+
|
| 365 |
+
def tiled_decode(self, z):
|
| 366 |
+
if self.parallel_tiling: # Fast online decoding for large videos
|
| 367 |
+
state = get_parallel_state()
|
| 368 |
+
sp_rank = state["sp_rank"]
|
| 369 |
+
sp_size = state["sp_size"]
|
| 370 |
+
else:
|
| 371 |
+
sp_rank, sp_size = 0, 1
|
| 372 |
+
|
| 373 |
+
height, width = (
|
| 374 |
+
z.shape[-2] * self.vae_ratio,
|
| 375 |
+
z.shape[-1] * self.vae_ratio,
|
| 376 |
+
)
|
| 377 |
+
y_idx, y_len, y_overlap = self.split_tiles(height, True)
|
| 378 |
+
x_idx, x_len, x_overlap = self.split_tiles(width, True)
|
| 379 |
+
|
| 380 |
+
i_max, j_max = len(y_idx), len(x_idx)
|
| 381 |
+
num_tiles = i_max * j_max
|
| 382 |
+
|
| 383 |
+
z_tiles = []
|
| 384 |
+
for i, (i_pos, i_len) in enumerate(zip(y_idx, y_len)):
|
| 385 |
+
i_pos, i_len = (
|
| 386 |
+
i_pos // self.vae_ratio,
|
| 387 |
+
i_len // self.vae_ratio,
|
| 388 |
+
)
|
| 389 |
+
for j, (j_pos, j_len) in enumerate(zip(x_idx, x_len)):
|
| 390 |
+
j_pos, j_len = (j_pos // self.vae_ratio, j_len // self.vae_ratio)
|
| 391 |
+
tile = z[..., i_pos : i_pos + i_len, j_pos : j_pos + j_len]
|
| 392 |
+
z_tiles.append(tile)
|
| 393 |
+
|
| 394 |
+
local_tile_indices = self._local_tile_indices(num_tiles, sp_rank, sp_size)
|
| 395 |
+
stack_tiling = self.stack_tiling and not (
|
| 396 |
+
self.training and getattr(self.decoder, "mask_enabled", False)
|
| 397 |
+
)
|
| 398 |
+
decoded_tasks = self._run_tile_tasks(
|
| 399 |
+
z_tiles, local_tile_indices, self.decode, stack_tiling
|
| 400 |
+
)
|
| 401 |
+
|
| 402 |
+
if sp_size > 1:
|
| 403 |
+
dist.barrier(group=get_parallel_state()["sp_process_group"])
|
| 404 |
+
all_decoded = self._all_gather_tiled_results(decoded_tasks, num_tiles)
|
| 405 |
+
else:
|
| 406 |
+
all_decoded = decoded_tasks
|
| 407 |
+
|
| 408 |
+
|
| 409 |
+
rows = [[None for _ in range(j_max)] for _ in range(i_max)]
|
| 410 |
+
for idx, decoded in enumerate(all_decoded):
|
| 411 |
+
i, j = idx // j_max, idx % j_max
|
| 412 |
+
rows[i][j] = decoded.to(z.device)
|
| 413 |
+
|
| 414 |
+
result_rows = []
|
| 415 |
+
for i, row in enumerate(rows):
|
| 416 |
+
result_row = []
|
| 417 |
+
for j, tile in enumerate(row):
|
| 418 |
+
if i > 0:
|
| 419 |
+
tile = self.blend(rows[i - 1][j], tile, y_overlap[i - 1], dim=-2)
|
| 420 |
+
if j > 0:
|
| 421 |
+
tile = self.blend(row[j - 1], tile, x_overlap[j - 1], dim=-1)
|
| 422 |
+
if i < len(rows) - 1:
|
| 423 |
+
tile = tile[..., : -y_overlap[i], :]
|
| 424 |
+
if j < len(row) - 1:
|
| 425 |
+
tile = tile[..., :, : -x_overlap[j]]
|
| 426 |
+
result_row.append(tile)
|
| 427 |
+
result_rows.append(torch.cat(result_row, dim=-1))
|
| 428 |
+
dec = torch.cat(result_rows, dim=-2)
|
| 429 |
+
return dec
|
| 430 |
+
|
| 431 |
+
def _adaptive_encode(self, x):
|
| 432 |
+
if self.encoder_tiling:
|
| 433 |
+
return self.tiled_encode(x)
|
| 434 |
+
else:
|
| 435 |
+
return self.encode(x)
|
| 436 |
+
|
| 437 |
+
def _adaptive_decode(self, z):
|
| 438 |
+
if self.decoder_tiling:
|
| 439 |
+
return self.tiled_decode(z)
|
| 440 |
+
else:
|
| 441 |
+
return self.decode(z)
|
| 442 |
+
|
| 443 |
+
def trim_code(self, z, target_codes):
|
| 444 |
+
if target_codes < z.shape[2]:
|
| 445 |
+
if self.causal_encoder:
|
| 446 |
+
z = z[:, :, -target_codes:, :, :]
|
| 447 |
+
else:
|
| 448 |
+
start_frame = (z.shape[2] - target_codes) // 2
|
| 449 |
+
z = z[:, :, start_frame : start_frame + target_codes, :, :]
|
| 450 |
+
return z
|
| 451 |
+
|
| 452 |
+
def trim_output(self, dec, target_frames):
|
| 453 |
+
if target_frames < dec.shape[2]:
|
| 454 |
+
if self.causal_encoder: # This is defined by encoder, not decoder
|
| 455 |
+
dec = dec[:, :, -target_frames:, :, :]
|
| 456 |
+
else:
|
| 457 |
+
start_frame = (dec.shape[2] - target_frames) // 2
|
| 458 |
+
dec = dec[:, :, start_frame : start_frame + target_frames, :, :]
|
| 459 |
+
return dec
|
| 460 |
+
|
| 461 |
+
def encode_temporal(self, x):
|
| 462 |
+
offset_frame = 1 if self.isolated_first_frame and self.frame_pre_padding == 0 else 0
|
| 463 |
+
|
| 464 |
+
if x.shape[2] % self.clip_length != offset_frame:
|
| 465 |
+
pad_size = (offset_frame - x.shape[2]) % self.clip_length
|
| 466 |
+
pad_frames = x[:, :, -1:].repeat(1, 1, pad_size, 1, 1)
|
| 467 |
+
x = torch.cat([x, pad_frames], dim=2)
|
| 468 |
+
|
| 469 |
+
num_chunks = (x.shape[2] - offset_frame) // self.clip_length
|
| 470 |
+
|
| 471 |
+
z_list = []
|
| 472 |
+
for i in range(num_chunks):
|
| 473 |
+
start_idx = i * self.clip_length + offset_frame
|
| 474 |
+
end_idx = (i + 1) * self.clip_length + offset_frame
|
| 475 |
+
clip_x = x[:, :, start_idx:end_idx, :, :]
|
| 476 |
+
|
| 477 |
+
if self.isolated_key_frame:
|
| 478 |
+
key_frame = clip_x[:, :, :1, :, :]
|
| 479 |
+
z_key = self._adaptive_encode(key_frame)
|
| 480 |
+
|
| 481 |
+
if clip_x.shape[2] > 1:
|
| 482 |
+
video_frames = clip_x[:, :, 1:, :, :]
|
| 483 |
+
z_video = self._adaptive_encode(video_frames)
|
| 484 |
+
z = torch.cat([z_key, z_video], dim=2)
|
| 485 |
+
else:
|
| 486 |
+
z = z_key
|
| 487 |
+
else:
|
| 488 |
+
z = self._adaptive_encode(clip_x)
|
| 489 |
+
|
| 490 |
+
z_list.append(z)
|
| 491 |
+
|
| 492 |
+
z = torch.cat(z_list, dim=2)
|
| 493 |
+
if self.token_drop > 0:
|
| 494 |
+
z = z[:, :, : -self.token_drop]
|
| 495 |
+
|
| 496 |
+
if self.isolated_first_frame:
|
| 497 |
+
input_first_frame = x[:, :, :1, :, :]
|
| 498 |
+
z_first_frame = self._adaptive_encode(input_first_frame)
|
| 499 |
+
|
| 500 |
+
if self.frame_pre_padding == 0:
|
| 501 |
+
z = torch.cat([z_first_frame, z], dim=2)
|
| 502 |
+
else:
|
| 503 |
+
z = torch.cat([z_first_frame, z[:, :, 1:, :, :]], dim=2)
|
| 504 |
+
|
| 505 |
+
if self.isolated_last_frame:
|
| 506 |
+
frame_num = x.shape[2]
|
| 507 |
+
last_frame_idx = frame_num - self.frame_drop + offset_frame
|
| 508 |
+
input_last_frame = x[:, :, last_frame_idx : last_frame_idx + 1, :, :]
|
| 509 |
+
z_last_frame = self._adaptive_encode(input_last_frame)
|
| 510 |
+
z = torch.cat([z, z_last_frame], dim=2)
|
| 511 |
+
|
| 512 |
+
return z
|
| 513 |
+
|
| 514 |
+
def _decode_temporal_pad_frames(self, z, pad_tokens):
|
| 515 |
+
if pad_tokens <= 0:
|
| 516 |
+
return 0
|
| 517 |
+
intra_tail = self.clip_length % self.vae_ratio_t
|
| 518 |
+
if intra_tail == 0:
|
| 519 |
+
return int(pad_tokens) * int(self.vae_ratio_t)
|
| 520 |
+
|
| 521 |
+
z_len_before_pad = z.shape[2] - pad_tokens
|
| 522 |
+
return sum(
|
| 523 |
+
(
|
| 524 |
+
intra_tail
|
| 525 |
+
if (z_len_before_pad + k) % self.tokens_chunk_size == 0
|
| 526 |
+
else self.vae_ratio_t
|
| 527 |
+
)
|
| 528 |
+
for k in range(pad_tokens)
|
| 529 |
+
)
|
| 530 |
+
|
| 531 |
+
def _decode_temporal_output_frame_plan(self, z, z_head, z_tail, num_chunks, pad_tokens):
|
| 532 |
+
chunk_dec = self.tokens_chunk_size * self.vae_ratio_t
|
| 533 |
+
split_count = int(self.token_drop > 0) + 1
|
| 534 |
+
total_frames = 0
|
| 535 |
+
final_overlap_frames = 0
|
| 536 |
+
|
| 537 |
+
if z_head is not None:
|
| 538 |
+
total_frames += 1
|
| 539 |
+
|
| 540 |
+
for i in range(num_chunks):
|
| 541 |
+
t_start_idx = i * self.tokens_chunk_size
|
| 542 |
+
t_end_idx = t_start_idx + self.tokens_chunk_size + self.token_overlap
|
| 543 |
+
clip_token_len = max(0, min(t_end_idx, z.shape[2]) - min(t_start_idx, z.shape[2]))
|
| 544 |
+
if i == 0 and z_head is not None:
|
| 545 |
+
clip_token_len += z_head.shape[2]
|
| 546 |
+
if i == num_chunks - 1 and z_tail is not None:
|
| 547 |
+
clip_token_len += z_tail.shape[2]
|
| 548 |
+
|
| 549 |
+
clip_frame_len = clip_token_len * self.vae_ratio_t
|
| 550 |
+
if i == 0 and z_head is not None:
|
| 551 |
+
clip_frame_len = max(0, clip_frame_len - self.vae_ratio_t)
|
| 552 |
+
if i == num_chunks - 1 and z_tail is not None:
|
| 553 |
+
clip_frame_len = max(0, clip_frame_len - self.vae_ratio_t)
|
| 554 |
+
|
| 555 |
+
for j in range(split_count):
|
| 556 |
+
f_start_idx = j * chunk_dec
|
| 557 |
+
f_end_idx = min(f_start_idx + chunk_dec, clip_frame_len)
|
| 558 |
+
chunk_frames = max(0, f_end_idx - f_start_idx - self.frame_pre_padding)
|
| 559 |
+
if j == 0:
|
| 560 |
+
total_frames += chunk_frames
|
| 561 |
+
else:
|
| 562 |
+
final_overlap_frames = chunk_frames
|
| 563 |
+
|
| 564 |
+
total_frames += final_overlap_frames
|
| 565 |
+
if z_tail is not None:
|
| 566 |
+
total_frames += 1
|
| 567 |
+
|
| 568 |
+
pad_frames = self._decode_temporal_pad_frames(z, pad_tokens)
|
| 569 |
+
return int(total_frames), int(pad_frames), int(total_frames - pad_frames)
|
| 570 |
+
|
| 571 |
+
def _decode_temporal_streaming(self, z, z_head, z_tail, num_chunks, pad_tokens, temporal_cat_dtype):
|
| 572 |
+
total_frames, pad_frames, output_frames = self._decode_temporal_output_frame_plan(
|
| 573 |
+
z, z_head, z_tail, num_chunks, pad_tokens
|
| 574 |
+
)
|
| 575 |
+
if output_frames <= 0:
|
| 576 |
+
raise ValueError(
|
| 577 |
+
f"decode_temporal streaming planned non-positive output_frames={output_frames} "
|
| 578 |
+
f"total_frames={total_frames} pad_frames={pad_frames}"
|
| 579 |
+
)
|
| 580 |
+
|
| 581 |
+
|
| 582 |
+
chunk_dec = self.tokens_chunk_size * self.vae_ratio_t
|
| 583 |
+
split_count = int(self.token_drop > 0) + 1
|
| 584 |
+
dec = None
|
| 585 |
+
dec_overlap = None
|
| 586 |
+
write_pos = 0
|
| 587 |
+
logical_frames = 0
|
| 588 |
+
dropped_frames = 0
|
| 589 |
+
decoded_count = 0
|
| 590 |
+
|
| 591 |
+
def write_part(part):
|
| 592 |
+
nonlocal dec, write_pos, logical_frames, dropped_frames
|
| 593 |
+
part_frames = int(part.shape[2])
|
| 594 |
+
if part_frames <= 0:
|
| 595 |
+
return
|
| 596 |
+
logical_frames += part_frames
|
| 597 |
+
if dec is None:
|
| 598 |
+
out_shape = list(part.shape)
|
| 599 |
+
out_shape[2] = output_frames
|
| 600 |
+
dec = torch.empty(out_shape, dtype=part.dtype, device=part.device)
|
| 601 |
+
|
| 602 |
+
remaining = int(dec.shape[2]) - write_pos
|
| 603 |
+
copy_frames = min(part_frames, max(0, remaining))
|
| 604 |
+
if copy_frames > 0:
|
| 605 |
+
dec[:, :, write_pos : write_pos + copy_frames, :, :].copy_(
|
| 606 |
+
part[:, :, :copy_frames, :, :]
|
| 607 |
+
)
|
| 608 |
+
write_pos += copy_frames
|
| 609 |
+
dropped_frames += part_frames - copy_frames
|
| 610 |
+
|
| 611 |
+
for i in range(num_chunks):
|
| 612 |
+
t_start_idx = i * self.tokens_chunk_size
|
| 613 |
+
t_end_idx = t_start_idx + self.tokens_chunk_size + self.token_overlap
|
| 614 |
+
clip_z = z[:, :, t_start_idx:t_end_idx, :, :]
|
| 615 |
+
|
| 616 |
+
if i == 0 and z_head is not None:
|
| 617 |
+
clip_z = torch.cat([z_head, clip_z], dim=2)
|
| 618 |
+
|
| 619 |
+
if i == num_chunks - 1 and z_tail is not None:
|
| 620 |
+
clip_z = torch.cat([clip_z, z_tail], dim=2)
|
| 621 |
+
|
| 622 |
+
clip_dec = self._adaptive_decode(clip_z)
|
| 623 |
+
decoded_count += 1
|
| 624 |
+
if temporal_cat_dtype is not None and clip_dec.dtype != temporal_cat_dtype:
|
| 625 |
+
clip_dec = clip_dec.to(temporal_cat_dtype)
|
| 626 |
+
if clip_dec.device != z.device:
|
| 627 |
+
clip_dec = clip_dec.to(z.device)
|
| 628 |
+
|
| 629 |
+
|
| 630 |
+
dec_tail = None
|
| 631 |
+
if i == 0 and z_head is not None:
|
| 632 |
+
write_part(clip_dec[:, :, self.vae_ratio_t - 1 : self.vae_ratio_t, :, :])
|
| 633 |
+
clip_dec = clip_dec[:, :, self.vae_ratio_t :, :, :]
|
| 634 |
+
|
| 635 |
+
if i == num_chunks - 1 and z_tail is not None:
|
| 636 |
+
dec_tail = clip_dec[:, :, -1:, :, :]
|
| 637 |
+
clip_dec = clip_dec[:, :, : -self.vae_ratio_t, :, :]
|
| 638 |
+
|
| 639 |
+
for j in range(split_count):
|
| 640 |
+
f_start_idx = j * chunk_dec
|
| 641 |
+
f_end_idx = min(f_start_idx + chunk_dec, clip_dec.shape[2])
|
| 642 |
+
clip_dec_chunk = clip_dec[:, :, f_start_idx:f_end_idx, :, :]
|
| 643 |
+
clip_dec_chunk = clip_dec_chunk[:, :, self.frame_pre_padding :, :, :]
|
| 644 |
+
|
| 645 |
+
if j == 0:
|
| 646 |
+
if dec_overlap is not None:
|
| 647 |
+
clip_dec_chunk = self.blend(
|
| 648 |
+
dec_overlap, clip_dec_chunk, self.frame_overlap, dim=-3
|
| 649 |
+
)
|
| 650 |
+
dec_overlap = None
|
| 651 |
+
write_part(clip_dec_chunk)
|
| 652 |
+
else:
|
| 653 |
+
# Break the view's reference to the full decoded clip so earlier
|
| 654 |
+
# temporal chunks can be released before the final output exists.
|
| 655 |
+
dec_overlap = clip_dec_chunk.contiguous()
|
| 656 |
+
|
| 657 |
+
if i == num_chunks - 1:
|
| 658 |
+
if dec_overlap is not None:
|
| 659 |
+
write_part(dec_overlap)
|
| 660 |
+
dec_overlap = None
|
| 661 |
+
if dec_tail is not None:
|
| 662 |
+
write_part(dec_tail)
|
| 663 |
+
|
| 664 |
+
del clip_dec, clip_z
|
| 665 |
+
|
| 666 |
+
if dec is None:
|
| 667 |
+
raise RuntimeError("decode_temporal streaming produced no output tensor")
|
| 668 |
+
if logical_frames != total_frames or dropped_frames != pad_frames or write_pos != output_frames:
|
| 669 |
+
raise RuntimeError(
|
| 670 |
+
"decode_temporal streaming frame plan mismatch: "
|
| 671 |
+
f"logical_frames={logical_frames} total_frames={total_frames} "
|
| 672 |
+
f"dropped_frames={dropped_frames} pad_frames={pad_frames} "
|
| 673 |
+
f"write_pos={write_pos} output_frames={output_frames}"
|
| 674 |
+
)
|
| 675 |
+
|
| 676 |
+
return dec
|
| 677 |
+
|
| 678 |
+
def decode_temporal(self, z):
|
| 679 |
+
chunk_dec = self.tokens_chunk_size * self.vae_ratio_t
|
| 680 |
+
|
| 681 |
+
isolated_token_num = 0
|
| 682 |
+
if self.isolated_first_frame and self.frame_pre_padding == 0:
|
| 683 |
+
isolated_token_num = isolated_token_num + 1
|
| 684 |
+
if self.isolated_last_frame:
|
| 685 |
+
isolated_token_num = isolated_token_num + 1
|
| 686 |
+
|
| 687 |
+
pseudo_total_tokens = z.shape[2] - isolated_token_num + self.token_drop
|
| 688 |
+
|
| 689 |
+
pad_tokens = 0
|
| 690 |
+
remainder = pseudo_total_tokens % self.tokens_chunk_size
|
| 691 |
+
if remainder != 0:
|
| 692 |
+
if self.training:
|
| 693 |
+
raise ValueError(f"Temporal token length {z.shape[2]} is wrong!")
|
| 694 |
+
else:
|
| 695 |
+
pad_tokens = self.tokens_chunk_size - remainder
|
| 696 |
+
pseudo_total_tokens = pseudo_total_tokens + pad_tokens
|
| 697 |
+
|
| 698 |
+
pseudo_num_chunks = pseudo_total_tokens // self.tokens_chunk_size
|
| 699 |
+
num_chunks = pseudo_num_chunks - int(self.token_drop > 0)
|
| 700 |
+
|
| 701 |
+
z_head = None
|
| 702 |
+
if self.isolated_first_frame and self.frame_pre_padding == 0:
|
| 703 |
+
z_head = z[:, :, :1, :, :]
|
| 704 |
+
z = z[:, :, 1:, :, :]
|
| 705 |
+
|
| 706 |
+
z_tail = None
|
| 707 |
+
if self.isolated_last_frame:
|
| 708 |
+
z_tail = z[:, :, -1:, :, :]
|
| 709 |
+
z = z[:, :, :-1, :, :]
|
| 710 |
+
|
| 711 |
+
if pad_tokens > 0:
|
| 712 |
+
pad_z = z[:, :, -1:, :, :].repeat(1, 1, pad_tokens, 1, 1)
|
| 713 |
+
z = torch.cat([z, pad_z], dim=2)
|
| 714 |
+
|
| 715 |
+
temporal_cat_dtype = _resolve_temporal_cat_dtype()
|
| 716 |
+
if not self.training and _resolve_temporal_stream_cat():
|
| 717 |
+
return self._decode_temporal_streaming(
|
| 718 |
+
z, z_head, z_tail, num_chunks, pad_tokens, temporal_cat_dtype
|
| 719 |
+
)
|
| 720 |
+
|
| 721 |
+
decoded_tasks = []
|
| 722 |
+
for i in range(num_chunks):
|
| 723 |
+
t_start_idx = i * self.tokens_chunk_size
|
| 724 |
+
t_end_idx = t_start_idx + self.tokens_chunk_size + self.token_overlap
|
| 725 |
+
clip_z = z[:, :, t_start_idx:t_end_idx, :, :]
|
| 726 |
+
|
| 727 |
+
if i == 0 and z_head is not None:
|
| 728 |
+
clip_z = torch.cat([z_head, clip_z], dim=2)
|
| 729 |
+
|
| 730 |
+
if i == num_chunks - 1 and z_tail is not None:
|
| 731 |
+
clip_z = torch.cat([clip_z, z_tail], dim=2)
|
| 732 |
+
|
| 733 |
+
clip_dec = self._adaptive_decode(clip_z)
|
| 734 |
+
if temporal_cat_dtype is not None and clip_dec.dtype != temporal_cat_dtype:
|
| 735 |
+
clip_dec = clip_dec.to(temporal_cat_dtype)
|
| 736 |
+
|
| 737 |
+
decoded_tasks.append((i, clip_dec))
|
| 738 |
+
|
| 739 |
+
clip_dec_list = [clip_dec.to(z.device) for _, clip_dec in decoded_tasks]
|
| 740 |
+
|
| 741 |
+
dec_list = []
|
| 742 |
+
dec_overlap = None
|
| 743 |
+
|
| 744 |
+
dec_head = None
|
| 745 |
+
if z_head is not None:
|
| 746 |
+
dec_head = clip_dec_list[0][:, :, self.vae_ratio_t - 1 : self.vae_ratio_t, :, :]
|
| 747 |
+
clip_dec_list[0] = clip_dec_list[0][:, :, self.vae_ratio_t :, :, :]
|
| 748 |
+
|
| 749 |
+
dec_tail = None
|
| 750 |
+
if z_tail is not None:
|
| 751 |
+
dec_tail = clip_dec_list[-1][:, :, -1:, :, :]
|
| 752 |
+
clip_dec_list[-1] = clip_dec_list[-1][:, :, : -self.vae_ratio_t, :, :]
|
| 753 |
+
|
| 754 |
+
if dec_head is not None:
|
| 755 |
+
dec_list.append(dec_head)
|
| 756 |
+
|
| 757 |
+
for i in range(num_chunks):
|
| 758 |
+
for j in range(int(self.token_drop > 0) + 1):
|
| 759 |
+
clip_dec = clip_dec_list[i]
|
| 760 |
+
|
| 761 |
+
f_start_idx = j * chunk_dec
|
| 762 |
+
f_end_idx = min(f_start_idx + chunk_dec, clip_dec.shape[2])
|
| 763 |
+
clip_dec_chunk = clip_dec[:, :, f_start_idx:f_end_idx, :, :]
|
| 764 |
+
clip_dec_chunk = clip_dec_chunk[:, :, self.frame_pre_padding :, :, :]
|
| 765 |
+
|
| 766 |
+
if j == 0:
|
| 767 |
+
if dec_overlap is not None:
|
| 768 |
+
clip_dec_chunk = self.blend(
|
| 769 |
+
dec_overlap, clip_dec_chunk, self.frame_overlap, dim=-3
|
| 770 |
+
)
|
| 771 |
+
dec_list.append(clip_dec_chunk)
|
| 772 |
+
else:
|
| 773 |
+
dec_overlap = clip_dec_chunk
|
| 774 |
+
|
| 775 |
+
if dec_overlap is not None:
|
| 776 |
+
dec_list.append(dec_overlap)
|
| 777 |
+
|
| 778 |
+
if dec_tail is not None:
|
| 779 |
+
dec_list.append(dec_tail)
|
| 780 |
+
|
| 781 |
+
|
| 782 |
+
dec = torch.cat(dec_list, dim=2)
|
| 783 |
+
|
| 784 |
+
pad_frames = self._decode_temporal_pad_frames(z, pad_tokens)
|
| 785 |
+
if pad_frames > 0:
|
| 786 |
+
dec = dec[:, :, :-pad_frames, :, :]
|
| 787 |
+
|
| 788 |
+
return dec
|
| 789 |
+
|
| 790 |
+
def decode_base(self, z, frame_num=None, process_image=False):
|
| 791 |
+
if process_image or not self.use_3d_conv:
|
| 792 |
+
if not self.use_3d_conv and z.ndim == 5:
|
| 793 |
+
z = z.squeeze(2)
|
| 794 |
+
|
| 795 |
+
recon = self._adaptive_decode(z)
|
| 796 |
+
else:
|
| 797 |
+
recon = self.decode_temporal(z)
|
| 798 |
+
|
| 799 |
+
if self.use_3d_conv:
|
| 800 |
+
if frame_num is not None:
|
| 801 |
+
target_frames = frame_num
|
| 802 |
+
else:
|
| 803 |
+
target_frames = recon.shape[2]
|
| 804 |
+
|
| 805 |
+
recon = self.trim_output(recon, target_frames)
|
| 806 |
+
if process_image:
|
| 807 |
+
recon = recon.squeeze(2)
|
| 808 |
+
|
| 809 |
+
return recon
|
| 810 |
+
|
| 811 |
+
#########################################################
|
| 812 |
+
# freeze_scope is retained from the training codebase: in this
|
| 813 |
+
# inference-only bundle (self.training is always False) it simply
|
| 814 |
+
# provides the no_grad() context used by encode()/decode().
|
| 815 |
+
#########################################################
|
| 816 |
+
|
| 817 |
+
|
| 818 |
+
def freeze_scope(self, module_name):
|
| 819 |
+
if not self.training:
|
| 820 |
+
return torch.no_grad()
|
| 821 |
+
|
| 822 |
+
if_freeze = module_name in self.fix_modules
|
| 823 |
+
if if_freeze:
|
| 824 |
+
return torch.no_grad()
|
| 825 |
+
else:
|
| 826 |
+
return nullcontext()
|
| 827 |
+
|
| 828 |
+
|
| 829 |
+
|
| 830 |
+
|
| 831 |
+
|
| 832 |
+
|
| 833 |
+
|
| 834 |
+
|
| 835 |
+
#########################################################
|
| 836 |
+
# following methods are for inference
|
| 837 |
+
#########################################################
|
| 838 |
+
|
| 839 |
+
@torch.no_grad()
|
| 840 |
+
def encode_images(
|
| 841 |
+
self,
|
| 842 |
+
images: Union[List[np.ndarray], List[torch.Tensor]],
|
| 843 |
+
transform_input: bool = False,
|
| 844 |
+
use_fp16_latent: bool = False,
|
| 845 |
+
verbose: bool = False,
|
| 846 |
+
) -> List[torch.Tensor]:
|
| 847 |
+
"""encode images into latents
|
| 848 |
+
|
| 849 |
+
Args:
|
| 850 |
+
images (Union[List[np.ndarray], List[torch.Tensor]]):
|
| 851 |
+
List of images, single input will be wrapped in a list.
|
| 852 |
+
If input is a list of np.ndarray, it should be in shape B * (H, W, 3), dtype uint8.
|
| 853 |
+
If input is a list of torch.Tensor, it should be in shape B * (3, H, W), dtype float32.
|
| 854 |
+
transform_input (bool, optional):
|
| 855 |
+
Whether to transform input using ImageNet std/mean. Defaults to False.
|
| 856 |
+
If input is a list of np.ndarray, it will always be set to True.
|
| 857 |
+
use_fp16_latent (bool, optional):
|
| 858 |
+
Whether to use fp16 latent. Defaults to False.
|
| 859 |
+
verbose (bool, optional):
|
| 860 |
+
Whether to print debug information. Defaults to False.
|
| 861 |
+
|
| 862 |
+
Returns:
|
| 863 |
+
List[torch.Tensor]:
|
| 864 |
+
List of image latents.
|
| 865 |
+
If self.use_3d_conv is True, it should be in shape B * (D, 1, H', W').
|
| 866 |
+
Otherwise, it should be in shape B * (D, H', W').
|
| 867 |
+
"""
|
| 868 |
+
|
| 869 |
+
images = self.processor._ensure_list(images)
|
| 870 |
+
|
| 871 |
+
if isinstance(images[0], Image.Image):
|
| 872 |
+
images = [np.array(image) for image in images]
|
| 873 |
+
|
| 874 |
+
if isinstance(images[0], np.ndarray):
|
| 875 |
+
device = next(self.parameters()).device
|
| 876 |
+
images = self.processor.convert_numpy_to_tensor(images, device)
|
| 877 |
+
images = torch.split(images, 1, dim=0)
|
| 878 |
+
transform_input = True
|
| 879 |
+
|
| 880 |
+
if transform_input:
|
| 881 |
+
images = [
|
| 882 |
+
image.unsqueeze(0) if image.ndim == 3 else image for image in images
|
| 883 |
+
]
|
| 884 |
+
images = [self.processor.transform_tensor(image) for image in images]
|
| 885 |
+
|
| 886 |
+
prepared = []
|
| 887 |
+
for image_tensor in images:
|
| 888 |
+
if image_tensor.ndim == 3:
|
| 889 |
+
image_tensor = image_tensor.unsqueeze(0)
|
| 890 |
+
_, _, h, w = image_tensor.shape
|
| 891 |
+
new_h, new_w = self.processor._align_to_total_patch_size(h, w)
|
| 892 |
+
image_tensor = self.processor._crop_to_align(image_tensor, new_h, new_w)
|
| 893 |
+
prepared.append(image_tensor)
|
| 894 |
+
|
| 895 |
+
if len(prepared) > 1 and len(set(t.shape for t in prepared)) == 1:
|
| 896 |
+
stacked = torch.cat(prepared, dim=0)
|
| 897 |
+
if verbose:
|
| 898 |
+
logger.info(f"batch encode input shape {tuple(stacked.shape)}")
|
| 899 |
+
all_latents = self.encode_base(stacked, True)
|
| 900 |
+
image_latents = [all_latents[i].contiguous() for i in range(all_latents.shape[0])]
|
| 901 |
+
else:
|
| 902 |
+
image_latents = []
|
| 903 |
+
for image_tensor in prepared:
|
| 904 |
+
if verbose:
|
| 905 |
+
logger.info(f"input shape {tuple(image_tensor.shape)}")
|
| 906 |
+
image_latent = self.encode_base(image_tensor, True)
|
| 907 |
+
image_latents.append(image_latent.squeeze(0).contiguous())
|
| 908 |
+
|
| 909 |
+
if use_fp16_latent:
|
| 910 |
+
image_latents = [lat.to(torch.float16) for lat in image_latents]
|
| 911 |
+
|
| 912 |
+
if verbose:
|
| 913 |
+
for lat in image_latents:
|
| 914 |
+
logger.info(f"image latent shape {tuple(lat.shape)}")
|
| 915 |
+
|
| 916 |
+
return image_latents
|
| 917 |
+
|
| 918 |
+
@torch.no_grad()
|
| 919 |
+
def encode_videos(
|
| 920 |
+
self,
|
| 921 |
+
videos: Union[List[np.ndarray], List[torch.Tensor]],
|
| 922 |
+
transform_input: bool = False,
|
| 923 |
+
use_fp16_latent: bool = False,
|
| 924 |
+
verbose: bool = False,
|
| 925 |
+
encode_prefix: bool = False,
|
| 926 |
+
) -> List[torch.Tensor]:
|
| 927 |
+
"""encode videos into latents
|
| 928 |
+
|
| 929 |
+
Args:
|
| 930 |
+
videos (Union[List[np.ndarray], List[torch.Tensor]]):
|
| 931 |
+
List of videos, single input will be wrapped in a list.
|
| 932 |
+
If input is a list of np.ndarray, it should be in shape B * (T, H, W, 3), dtype uint8.
|
| 933 |
+
If input is a list of torch.Tensor, it should be in shape B * (3, T, H, W), dtype float32.
|
| 934 |
+
transform_input (bool, optional):
|
| 935 |
+
Whether to transform input using ImageNet std/mean. Defaults to False.
|
| 936 |
+
If input is a list of np.ndarray, it will always be set to True.
|
| 937 |
+
use_fp16_latent (bool, optional):
|
| 938 |
+
Whether to use fp16 latent. Defaults to False.
|
| 939 |
+
verbose (bool, optional):
|
| 940 |
+
Whether to print debug information. Defaults to False.
|
| 941 |
+
encode_prefix (bool, optional):
|
| 942 |
+
Continuation (prefix) mode: prepend normalized
|
| 943 |
+
black frames to token alignment, append black frames to chunk
|
| 944 |
+
alignment, encode with token_drop disabled, then discard only
|
| 945 |
+
the trailing padding tokens. Returns both latents and leading
|
| 946 |
+
pad-frame counts. Defaults to False.
|
| 947 |
+
|
| 948 |
+
Returns:
|
| 949 |
+
List[torch.Tensor]:
|
| 950 |
+
List of video latents, shape B * (D, T', H', W').
|
| 951 |
+
With encode_prefix=True, returns
|
| 952 |
+
(List[torch.Tensor], List[int]).
|
| 953 |
+
"""
|
| 954 |
+
|
| 955 |
+
videos = self.processor._ensure_list(videos)
|
| 956 |
+
|
| 957 |
+
if isinstance(videos[0], np.ndarray):
|
| 958 |
+
device = next(self.parameters()).device
|
| 959 |
+
videos = [self.processor.convert_numpy_to_tensor(video, device) for video in videos]
|
| 960 |
+
transform_input = True
|
| 961 |
+
|
| 962 |
+
if transform_input:
|
| 963 |
+
videos = [self.processor.transform_tensor(video) for video in videos]
|
| 964 |
+
videos = [video.transpose(0, 1) for video in videos]
|
| 965 |
+
|
| 966 |
+
if encode_prefix:
|
| 967 |
+
if self.isolated_last_frame:
|
| 968 |
+
raise ValueError(
|
| 969 |
+
"encode_prefix does not support isolated_last_frame"
|
| 970 |
+
)
|
| 971 |
+
|
| 972 |
+
video_latents = []
|
| 973 |
+
prefix_pad_frames = []
|
| 974 |
+
for video in videos:
|
| 975 |
+
if video.ndim == 4:
|
| 976 |
+
video = video.unsqueeze(0)
|
| 977 |
+
_, _, _, h, w = video.shape
|
| 978 |
+
new_h, new_w = self.processor._align_to_total_patch_size(h, w)
|
| 979 |
+
video = self.processor._crop_to_align(
|
| 980 |
+
video, new_h, new_w, is_video=True
|
| 981 |
+
)
|
| 982 |
+
|
| 983 |
+
model_alignment = (
|
| 984 |
+
self.token_drop,
|
| 985 |
+
self.frame_drop,
|
| 986 |
+
self.token_overlap,
|
| 987 |
+
self.frame_overlap,
|
| 988 |
+
)
|
| 989 |
+
processor_alignment = (
|
| 990 |
+
self.processor.token_overlap,
|
| 991 |
+
self.processor.frame_overlap,
|
| 992 |
+
)
|
| 993 |
+
self.token_drop = 0
|
| 994 |
+
self.frame_drop = 0
|
| 995 |
+
self.token_overlap = 0
|
| 996 |
+
self.frame_overlap = 0
|
| 997 |
+
self.processor.token_overlap = 0
|
| 998 |
+
self.processor.frame_overlap = 0
|
| 999 |
+
try:
|
| 1000 |
+
orig_frames = video.shape[2]
|
| 1001 |
+
leading, trailing, drop_tokens = (
|
| 1002 |
+
self.processor.align_video_length_2pass(orig_frames)
|
| 1003 |
+
)
|
| 1004 |
+
_, _, _, cropped_h, cropped_w = video.shape
|
| 1005 |
+
if leading > 0:
|
| 1006 |
+
black = self.processor.transform(
|
| 1007 |
+
video.new_zeros(leading, 3, cropped_h, cropped_w)
|
| 1008 |
+
)
|
| 1009 |
+
black = black.unsqueeze(0).permute(0, 2, 1, 3, 4)
|
| 1010 |
+
video = torch.cat([black, video], dim=2)
|
| 1011 |
+
if trailing > 0:
|
| 1012 |
+
black = self.processor.transform(
|
| 1013 |
+
video.new_zeros(trailing, 3, cropped_h, cropped_w)
|
| 1014 |
+
)
|
| 1015 |
+
black = black.unsqueeze(0).permute(0, 2, 1, 3, 4)
|
| 1016 |
+
video = torch.cat([video, black], dim=2)
|
| 1017 |
+
|
| 1018 |
+
if verbose:
|
| 1019 |
+
logger.info(
|
| 1020 |
+
f"[encode_prefix] {orig_frames} frames -> "
|
| 1021 |
+
f"pad leading={leading}, trailing={trailing} -> "
|
| 1022 |
+
f"{video.shape[2]} frames"
|
| 1023 |
+
)
|
| 1024 |
+
|
| 1025 |
+
video_latent = self.encode_base(video, False)
|
| 1026 |
+
if drop_tokens > 0:
|
| 1027 |
+
video_latent = video_latent[:, :, :-drop_tokens, :, :]
|
| 1028 |
+
prefix_pad_frames.append(leading)
|
| 1029 |
+
finally:
|
| 1030 |
+
(
|
| 1031 |
+
self.token_drop,
|
| 1032 |
+
self.frame_drop,
|
| 1033 |
+
self.token_overlap,
|
| 1034 |
+
self.frame_overlap,
|
| 1035 |
+
) = model_alignment
|
| 1036 |
+
(
|
| 1037 |
+
self.processor.token_overlap,
|
| 1038 |
+
self.processor.frame_overlap,
|
| 1039 |
+
) = processor_alignment
|
| 1040 |
+
|
| 1041 |
+
video_latents.append(video_latent.squeeze(0).contiguous())
|
| 1042 |
+
|
| 1043 |
+
if use_fp16_latent:
|
| 1044 |
+
video_latents = [lat.to(torch.float16) for lat in video_latents]
|
| 1045 |
+
if verbose:
|
| 1046 |
+
for latent in video_latents:
|
| 1047 |
+
logger.info(f"video latent shape {tuple(latent.shape)}")
|
| 1048 |
+
return video_latents, prefix_pad_frames
|
| 1049 |
+
|
| 1050 |
+
prepared = []
|
| 1051 |
+
for video in videos:
|
| 1052 |
+
if video.ndim == 4:
|
| 1053 |
+
video = video.unsqueeze(0)
|
| 1054 |
+
used_frame_length = self.processor.get_suitable_video_length(video.shape[2], verbose)
|
| 1055 |
+
_, _, _, h, w = video.shape
|
| 1056 |
+
new_h, new_w = self.processor._align_to_total_patch_size(h, w)
|
| 1057 |
+
video = video[:, :, :used_frame_length, :, :]
|
| 1058 |
+
video = self.processor._crop_to_align(video, new_h, new_w, is_video=True)
|
| 1059 |
+
prepared.append(video)
|
| 1060 |
+
|
| 1061 |
+
if len(prepared) > 1 and len(set(t.shape for t in prepared)) == 1:
|
| 1062 |
+
stacked = torch.cat(prepared, dim=0)
|
| 1063 |
+
if verbose:
|
| 1064 |
+
logger.info(f"batch encode input shape {tuple(stacked.shape)}")
|
| 1065 |
+
all_latents = self.encode_base(stacked, False)
|
| 1066 |
+
video_latents = [all_latents[i].contiguous() for i in range(all_latents.shape[0])]
|
| 1067 |
+
else:
|
| 1068 |
+
video_latents = []
|
| 1069 |
+
for video in prepared:
|
| 1070 |
+
if verbose:
|
| 1071 |
+
logger.info(f"input shape {tuple(video.shape)}")
|
| 1072 |
+
video_latent = self.encode_base(video, False)
|
| 1073 |
+
video_latents.append(video_latent.squeeze(0).contiguous())
|
| 1074 |
+
|
| 1075 |
+
if use_fp16_latent:
|
| 1076 |
+
video_latents = [lat.to(torch.float16) for lat in video_latents]
|
| 1077 |
+
|
| 1078 |
+
if verbose:
|
| 1079 |
+
for lat in video_latents:
|
| 1080 |
+
logger.info(f"video latent shape {tuple(lat.shape)}")
|
| 1081 |
+
|
| 1082 |
+
return video_latents
|
| 1083 |
+
|
| 1084 |
+
|
| 1085 |
+
|
| 1086 |
+
|
| 1087 |
+
# ============================================================================
|
| 1088 |
+
# Legacy CNN VAE
|
| 1089 |
+
# ============================================================================
|
| 1090 |
+
|
| 1091 |
+
|
| 1092 |
+
class AutoencoderKLLegacy(AutoencoderKL):
|
| 1093 |
+
r"""
|
| 1094 |
+
A VAE model (legacy CNN-based) for encoding pixels into latents and decoding latent representations into pixels.
|
| 1095 |
+
"""
|
| 1096 |
+
|
| 1097 |
+
@register_to_config
|
| 1098 |
+
def __init__(
|
| 1099 |
+
self,
|
| 1100 |
+
in_channels=3,
|
| 1101 |
+
out_ch=3,
|
| 1102 |
+
ch=128,
|
| 1103 |
+
embed_dim=16,
|
| 1104 |
+
z_channels=16,
|
| 1105 |
+
use_3d_conv=False,
|
| 1106 |
+
# cnn vae
|
| 1107 |
+
zq_ch_encoder=None,
|
| 1108 |
+
zq_ch_decoder=None,
|
| 1109 |
+
num_res_blocks=2,
|
| 1110 |
+
num_res_blocks_decoder=None,
|
| 1111 |
+
ch_mult=[1, 2, 2, 4, 4, 8],
|
| 1112 |
+
space_down=[2, 2, 2, 2, 1, 1],
|
| 1113 |
+
space_up=[1, 2, 2, 2, 2, 1],
|
| 1114 |
+
time_down=None,
|
| 1115 |
+
time_up=None,
|
| 1116 |
+
padding_mode="zeros",
|
| 1117 |
+
padding_mode_t=None,
|
| 1118 |
+
use_t_isolated_gn=False,
|
| 1119 |
+
causal_encoder=True,
|
| 1120 |
+
causal_decoder=True,
|
| 1121 |
+
use_vit_decoder=False,
|
| 1122 |
+
vit_decoder_kwargs=None,
|
| 1123 |
+
# stats
|
| 1124 |
+
shift_factor=0.0,
|
| 1125 |
+
scaling_factor=1.0,
|
| 1126 |
+
# pixel normalization
|
| 1127 |
+
pixel_norm_type="imagenet",
|
| 1128 |
+
# others
|
| 1129 |
+
**kwargs,
|
| 1130 |
+
):
|
| 1131 |
+
ModelMixin.__init__(self) # NOTE: avoid wrong @register_to_config
|
| 1132 |
+
|
| 1133 |
+
if not use_3d_conv or not use_vit_decoder:
|
| 1134 |
+
raise NotImplementedError(
|
| 1135 |
+
"this release only supports use_3d_conv=True with use_vit_decoder=True"
|
| 1136 |
+
)
|
| 1137 |
+
|
| 1138 |
+
self.transform = get_normalize_transform(pixel_norm_type)
|
| 1139 |
+
self.transform_rev = get_denormalize_transform(pixel_norm_type)
|
| 1140 |
+
|
| 1141 |
+
self.use_3d_conv = use_3d_conv
|
| 1142 |
+
self.causal_encoder = causal_encoder
|
| 1143 |
+
self.causal_decoder = causal_decoder
|
| 1144 |
+
self.slidedec = self.causal_encoder and not self.causal_decoder
|
| 1145 |
+
|
| 1146 |
+
# some registered parameters for simplicity
|
| 1147 |
+
self.vae_ratio = int(np.cumprod(space_down)[-1])
|
| 1148 |
+
self.vae_ratio_t = int(np.cumprod(time_down)[-1]) if time_down else 1
|
| 1149 |
+
self.config["vae_ratio"] = self.vae_ratio
|
| 1150 |
+
self.config["vae_ratio_t"] = self.vae_ratio_t
|
| 1151 |
+
|
| 1152 |
+
# some registered parameters for inference and training
|
| 1153 |
+
self.setup_forward(**kwargs)
|
| 1154 |
+
self.setup_training(**kwargs)
|
| 1155 |
+
|
| 1156 |
+
# init encoder
|
| 1157 |
+
encoder_config = {
|
| 1158 |
+
"double_z": True,
|
| 1159 |
+
"z_channels": z_channels,
|
| 1160 |
+
"zq_ch": zq_ch_encoder,
|
| 1161 |
+
"in_channels": in_channels,
|
| 1162 |
+
"ch": ch,
|
| 1163 |
+
"num_res_blocks": num_res_blocks,
|
| 1164 |
+
"ch_mult": ch_mult,
|
| 1165 |
+
"space_down": space_down,
|
| 1166 |
+
"time_down": time_down,
|
| 1167 |
+
"padding_mode": padding_mode,
|
| 1168 |
+
"padding_mode_t": padding_mode_t,
|
| 1169 |
+
"causal": causal_encoder,
|
| 1170 |
+
"use_t_isolated_gn": use_t_isolated_gn,
|
| 1171 |
+
}
|
| 1172 |
+
self.encoder = EncoderFCN3D(**encoder_config)
|
| 1173 |
+
|
| 1174 |
+
# init pointwise quant/post_quant conv
|
| 1175 |
+
self.quant_conv = nn.Conv3d(z_channels * 2, 2 * embed_dim, 1)
|
| 1176 |
+
self.post_quant_conv = nn.Conv3d(embed_dim, z_channels, 1)
|
| 1177 |
+
|
| 1178 |
+
self.use_vit_decoder = use_vit_decoder
|
| 1179 |
+
|
| 1180 |
+
# init decoder
|
| 1181 |
+
vit_kwargs = {
|
| 1182 |
+
"patch_size": self.vae_ratio,
|
| 1183 |
+
"in_channels": z_channels,
|
| 1184 |
+
"out_channels": out_ch,
|
| 1185 |
+
**(vit_decoder_kwargs or {}),
|
| 1186 |
+
}
|
| 1187 |
+
vit_kwargs.setdefault("patch_size_t", self.vae_ratio_t)
|
| 1188 |
+
vit_kwargs.setdefault("t_causal", causal_decoder)
|
| 1189 |
+
self.decoder = ViT3DDecoder(**vit_kwargs)
|
| 1190 |
+
|
| 1191 |
+
apply_spatial_parallel(self.encoder, self.encoder_parallel, self.chunk_dim)
|
| 1192 |
+
apply_spatial_parallel(self.decoder, self.decoder_parallel, self.chunk_dim)
|
| 1193 |
+
|
| 1194 |
+
for module in set(self.fix_modules + self.frozen_modules):
|
| 1195 |
+
self._freeze_nested_module(module)
|
| 1196 |
+
|
| 1197 |
+
self.gradient_checkpointing = False
|
| 1198 |
+
|
| 1199 |
+
def encode(self, x):
|
| 1200 |
+
if self.encoder_parallel:
|
| 1201 |
+
x = self.perform_input_slice(x, self.vae_ratio)
|
| 1202 |
+
|
| 1203 |
+
with self.freeze_scope("encoder"):
|
| 1204 |
+
h = self.encoder(x)
|
| 1205 |
+
|
| 1206 |
+
with self.freeze_scope("quant_conv"):
|
| 1207 |
+
moments = self.quant_conv(h)
|
| 1208 |
+
|
| 1209 |
+
if self.encoder_parallel:
|
| 1210 |
+
moments = self.perform_output_concat(moments)
|
| 1211 |
+
|
| 1212 |
+
return moments
|
| 1213 |
+
|
| 1214 |
+
def decode(self, z):
|
| 1215 |
+
if self.decoder_parallel and not self.use_vit_decoder:
|
| 1216 |
+
z = self.perform_input_slice(z)
|
| 1217 |
+
|
| 1218 |
+
with self.freeze_scope("post_quant_conv"):
|
| 1219 |
+
z2 = self.post_quant_conv(z)
|
| 1220 |
+
|
| 1221 |
+
with self.freeze_scope("decoder"):
|
| 1222 |
+
if self.use_vit_decoder:
|
| 1223 |
+
dec = self.decoder(z2)
|
| 1224 |
+
else:
|
| 1225 |
+
dec = self.decoder(z2, z)
|
| 1226 |
+
|
| 1227 |
+
if self.decoder_parallel and not self.use_vit_decoder:
|
| 1228 |
+
dec = self.perform_output_concat(dec)
|
| 1229 |
+
return dec
|
| 1230 |
+
|
| 1231 |
+
def encode_base(self, input, process_image=False):
|
| 1232 |
+
if self.use_3d_conv and input.ndim == 4:
|
| 1233 |
+
input = input.unsqueeze(2)
|
| 1234 |
+
|
| 1235 |
+
if process_image or not self.use_3d_conv:
|
| 1236 |
+
moments = self._adaptive_encode(input)
|
| 1237 |
+
else:
|
| 1238 |
+
moments = self.encode_temporal(input)
|
| 1239 |
+
|
| 1240 |
+
z = DiagonalGaussianDistribution(moments).sample()
|
| 1241 |
+
|
| 1242 |
+
if process_image and self.use_3d_conv:
|
| 1243 |
+
z = self.trim_code(z, 1)
|
| 1244 |
+
|
| 1245 |
+
return z
|
| 1246 |
+
|
| 1247 |
+
#########################################################
|
| 1248 |
+
# training-related knobs kept only for checkpoint/config compatibility
|
| 1249 |
+
#########################################################
|
| 1250 |
+
|
| 1251 |
+
def setup_training(self, **kwargs):
|
| 1252 |
+
self.fix_modules = kwargs.get("fix_modules", [])
|
| 1253 |
+
self.frozen_modules = kwargs.get("frozen_modules", [])
|
| 1254 |
+
|
| 1255 |
+
|
| 1256 |
+
|
| 1257 |
+
|
| 1258 |
+
|
FL2VA/video_vae/minimax_h3_video_vae.py
ADDED
|
@@ -0,0 +1,122 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 2 |
+
# Remote entry: self-contained MiniMax H3 visual VAE (3D CNN encoder + ViT3D decoder).
|
| 3 |
+
# Loaded via config.json:auto_map with trust_remote_code.
|
| 4 |
+
from __future__ import annotations
|
| 5 |
+
|
| 6 |
+
import json
|
| 7 |
+
from pathlib import Path
|
| 8 |
+
|
| 9 |
+
import safetensors.torch
|
| 10 |
+
import torch.nn as nn
|
| 11 |
+
|
| 12 |
+
# --- dependency manifest ---
|
| 13 |
+
# diffusers' dynamic-module loader only copies ONE level of relative
|
| 14 |
+
# imports into its cache; list every bundle module here so all files
|
| 15 |
+
# are copied, letting their own second-level imports resolve.
|
| 16 |
+
from .attention import Attention as _dep_attention # noqa: F401
|
| 17 |
+
from .base_module import FeedForward as _dep_base_module # noqa: F401
|
| 18 |
+
from .conv import SpatialParallelConv3d as _dep_conv # noqa: F401
|
| 19 |
+
from .flash import make_block_causal_mask_mod as _dep_flash # noqa: F401
|
| 20 |
+
from .func import create_token_ids as _dep_func # noqa: F401
|
| 21 |
+
from .klvae import AutoencoderKL as _dep_klvae # noqa: F401
|
| 22 |
+
from .norm import FusedGroupNorm3D as _dep_norm # noqa: F401
|
| 23 |
+
from .normalize import get_norm_constants as _dep_normalize # noqa: F401
|
| 24 |
+
from .parallel import get_parallel_state as _dep_parallel # noqa: F401
|
| 25 |
+
from .utils import apply_spatial_parallel as _dep_utils # noqa: F401
|
| 26 |
+
from .vae_cnn import EncoderFCN3D as _dep_vae_cnn # noqa: F401
|
| 27 |
+
from .vae_module import DiagonalGaussianDistribution as _dep_vae_module # noqa: F401
|
| 28 |
+
from .vae_processor import VAEProcessor as _dep_vae_processor # noqa: F401
|
| 29 |
+
from .vae_vit import ViTBase as _dep_vae_vit # noqa: F401
|
| 30 |
+
# --- end dependency manifest ---
|
| 31 |
+
|
| 32 |
+
from .klvae import AutoencoderKLLegacy
|
| 33 |
+
from .parallel import get_parallel_state
|
| 34 |
+
|
| 35 |
+
_SOURCE_CLASSES = {
|
| 36 |
+
"AutoencoderKLLegacy": AutoencoderKLLegacy,
|
| 37 |
+
}
|
| 38 |
+
|
| 39 |
+
|
| 40 |
+
def _ensure_vae_parallel_state() -> None:
|
| 41 |
+
"""Seed the bundled VAE parallel state for single-process inference."""
|
| 42 |
+
state = get_parallel_state()
|
| 43 |
+
if not isinstance(state, dict):
|
| 44 |
+
raise TypeError("get_parallel_state() must return a dict")
|
| 45 |
+
if state:
|
| 46 |
+
return
|
| 47 |
+
state.update(
|
| 48 |
+
{
|
| 49 |
+
"group_size": 1,
|
| 50 |
+
"group_rank": 0,
|
| 51 |
+
"local_process_group": None,
|
| 52 |
+
"sp_size": 1,
|
| 53 |
+
"sp_rank": 0,
|
| 54 |
+
"sp_enabled": False,
|
| 55 |
+
"sp_process_group": None,
|
| 56 |
+
"tp_size": 1,
|
| 57 |
+
"tp_rank": 0,
|
| 58 |
+
}
|
| 59 |
+
)
|
| 60 |
+
|
| 61 |
+
|
| 62 |
+
class MiniMaxH3VideoVAE(nn.Module):
|
| 63 |
+
def __init__(self, model: nn.Module) -> None:
|
| 64 |
+
super().__init__()
|
| 65 |
+
self.model = model
|
| 66 |
+
|
| 67 |
+
@classmethod
|
| 68 |
+
def from_pretrained(cls, pretrained_model_name_or_path: str, **kwargs):
|
| 69 |
+
component_dir = Path(pretrained_model_name_or_path)
|
| 70 |
+
with (component_dir / "config.json").open("r", encoding="utf-8") as f:
|
| 71 |
+
config = json.load(f)
|
| 72 |
+
source_path = component_dir / config["source_path"]
|
| 73 |
+
source_class_name = config["source_class_name"]
|
| 74 |
+
if source_class_name not in _SOURCE_CLASSES:
|
| 75 |
+
raise ValueError(
|
| 76 |
+
f"unsupported source_class_name {source_class_name!r}; "
|
| 77 |
+
f"bundled: {sorted(_SOURCE_CLASSES)}"
|
| 78 |
+
)
|
| 79 |
+
source_cls = _SOURCE_CLASSES[source_class_name]
|
| 80 |
+
if "source_safetensors_path" not in config:
|
| 81 |
+
raise ValueError(
|
| 82 |
+
"source_safetensors_path is required; pickle checkpoints are "
|
| 83 |
+
"not supported"
|
| 84 |
+
)
|
| 85 |
+
weights_path = source_path / config["source_safetensors_path"]
|
| 86 |
+
if not weights_path.is_file():
|
| 87 |
+
raise FileNotFoundError(f"source weights not found: {weights_path}")
|
| 88 |
+
if bool(config["vae_parallel_tiling"]):
|
| 89 |
+
_ensure_vae_parallel_state()
|
| 90 |
+
load_kwargs = {
|
| 91 |
+
"clip_length": int(config["vae_clip_length"]),
|
| 92 |
+
"token_drop": int(config["vae_token_drop"]),
|
| 93 |
+
"encoder_tiling": int(config["vae_encoder_tiling"]),
|
| 94 |
+
"decoder_tiling": int(config["vae_decoder_tiling"]),
|
| 95 |
+
"parallel_tiling": int(config["vae_parallel_tiling"]),
|
| 96 |
+
"tile_size": int(config["vae_tile_size"]),
|
| 97 |
+
"tile_overlap_min": int(config["vae_tile_overlap_min"]),
|
| 98 |
+
"encoder_parallel": int(config["vae_encoder_parallel"]),
|
| 99 |
+
"decoder_parallel": int(config["vae_decoder_parallel"]),
|
| 100 |
+
"chunk_dim": int(config["vae_chunk_dim"]),
|
| 101 |
+
}
|
| 102 |
+
# Mirror diffusers ModelMixin.from_pretrained instantiation semantics
|
| 103 |
+
# (config-driven init via from_config) but load the state dict from an
|
| 104 |
+
# explicitly named safetensors file instead of the diffusers default
|
| 105 |
+
# weight filename.
|
| 106 |
+
source_config = source_cls.load_config(str(source_path))
|
| 107 |
+
model, _unused = source_cls.from_config(
|
| 108 |
+
source_config, return_unused_kwargs=True, **load_kwargs
|
| 109 |
+
)
|
| 110 |
+
state_dict = safetensors.torch.load_file(str(weights_path))
|
| 111 |
+
model.load_state_dict(state_dict, strict=True)
|
| 112 |
+
model.eval()
|
| 113 |
+
return cls(model)
|
| 114 |
+
|
| 115 |
+
def forward(self, *args, **kwargs):
|
| 116 |
+
return self.model(*args, **kwargs)
|
| 117 |
+
|
| 118 |
+
def __getattr__(self, name: str):
|
| 119 |
+
try:
|
| 120 |
+
return super().__getattr__(name)
|
| 121 |
+
except AttributeError:
|
| 122 |
+
return getattr(self.model, name)
|
FL2VA/video_vae/norm.py
ADDED
|
@@ -0,0 +1,357 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 2 |
+
# Torch-native normalization for the MiniMax H3 visual VAE.
|
| 3 |
+
import math
|
| 4 |
+
import os
|
| 5 |
+
|
| 6 |
+
import torch
|
| 7 |
+
import torch.distributed as dist
|
| 8 |
+
import torch.nn as nn
|
| 9 |
+
import torch.nn.functional as F
|
| 10 |
+
|
| 11 |
+
from .conv import SpatialParallelConv3d
|
| 12 |
+
from .parallel import all_reduce, get_parallel_state
|
| 13 |
+
|
| 14 |
+
|
| 15 |
+
def _validate_activation(activation):
|
| 16 |
+
valid_activations = {"identity", "silu", "relu"}
|
| 17 |
+
if activation not in valid_activations:
|
| 18 |
+
raise ValueError(
|
| 19 |
+
f"Unsupported activation: {activation}. Supported: {valid_activations}"
|
| 20 |
+
)
|
| 21 |
+
|
| 22 |
+
|
| 23 |
+
def _apply_activation(x, activation):
|
| 24 |
+
_validate_activation(activation)
|
| 25 |
+
if activation == "identity":
|
| 26 |
+
return x
|
| 27 |
+
if activation == "silu":
|
| 28 |
+
return F.silu(x)
|
| 29 |
+
return F.relu(x)
|
| 30 |
+
|
| 31 |
+
|
| 32 |
+
def _merge_time_to_batch(x):
|
| 33 |
+
batch, channels, depth, height, width = x.shape
|
| 34 |
+
return (
|
| 35 |
+
x.permute(0, 2, 1, 3, 4)
|
| 36 |
+
.contiguous()
|
| 37 |
+
.view(batch * depth, channels, 1, height, width)
|
| 38 |
+
)
|
| 39 |
+
|
| 40 |
+
|
| 41 |
+
def _split_time_from_batch(x, batch):
|
| 42 |
+
batch_depth, channels, _, height, width = x.shape
|
| 43 |
+
depth = batch_depth // batch
|
| 44 |
+
return (
|
| 45 |
+
x.view(batch, depth, channels, height, width)
|
| 46 |
+
.permute(0, 2, 1, 3, 4)
|
| 47 |
+
.contiguous()
|
| 48 |
+
)
|
| 49 |
+
|
| 50 |
+
|
| 51 |
+
def fused_group_norm(x, num_groups, weight, bias, eps=1e-5, activation="silu"):
|
| 52 |
+
out = F.group_norm(x, num_groups, weight=weight, bias=bias, eps=eps)
|
| 53 |
+
return _apply_activation(out, activation)
|
| 54 |
+
|
| 55 |
+
|
| 56 |
+
def fused_spatial_norm(
|
| 57 |
+
f,
|
| 58 |
+
num_groups,
|
| 59 |
+
norm_weight,
|
| 60 |
+
norm_bias,
|
| 61 |
+
dynamic_scale,
|
| 62 |
+
dynamic_bias,
|
| 63 |
+
eps=1e-5,
|
| 64 |
+
activation="silu",
|
| 65 |
+
):
|
| 66 |
+
norm_f = F.group_norm(
|
| 67 |
+
f,
|
| 68 |
+
num_groups,
|
| 69 |
+
weight=norm_weight,
|
| 70 |
+
bias=norm_bias,
|
| 71 |
+
eps=eps,
|
| 72 |
+
)
|
| 73 |
+
out = norm_f * dynamic_scale + dynamic_bias
|
| 74 |
+
return _apply_activation(out, activation)
|
| 75 |
+
|
| 76 |
+
|
| 77 |
+
class DummyAffine(torch.nn.Module):
|
| 78 |
+
def __init__(self, num_channels, affine=True):
|
| 79 |
+
super().__init__()
|
| 80 |
+
if affine:
|
| 81 |
+
self.weight = torch.nn.Parameter(torch.ones(num_channels))
|
| 82 |
+
self.bias = torch.nn.Parameter(torch.zeros(num_channels))
|
| 83 |
+
else:
|
| 84 |
+
self.register_parameter("weight", None)
|
| 85 |
+
self.register_parameter("bias", None)
|
| 86 |
+
|
| 87 |
+
def forward(self, input):
|
| 88 |
+
if self.weight is None:
|
| 89 |
+
return input
|
| 90 |
+
shape = [1, -1] + [1] * (input.dim() - 2)
|
| 91 |
+
return input * self.weight.view(*shape) + self.bias.view(*shape)
|
| 92 |
+
|
| 93 |
+
|
| 94 |
+
class FusedGroupNorm3D(torch.nn.Module):
|
| 95 |
+
"""Compatibility wrapper implemented with native PyTorch ops."""
|
| 96 |
+
|
| 97 |
+
def __init__(
|
| 98 |
+
self,
|
| 99 |
+
num_groups,
|
| 100 |
+
num_channels,
|
| 101 |
+
eps=1e-5,
|
| 102 |
+
affine=True,
|
| 103 |
+
activation="silu",
|
| 104 |
+
cond_channels=None,
|
| 105 |
+
use_t_isolated_gn=False,
|
| 106 |
+
padding_mode="zeros",
|
| 107 |
+
padding_mode_t=None,
|
| 108 |
+
causal=True,
|
| 109 |
+
):
|
| 110 |
+
super().__init__()
|
| 111 |
+
_validate_activation(activation)
|
| 112 |
+
self.num_groups = num_groups
|
| 113 |
+
self.num_channels = num_channels
|
| 114 |
+
self.eps = eps
|
| 115 |
+
self.affine = affine
|
| 116 |
+
self.activation = activation
|
| 117 |
+
self.use_t_isolated_gn = use_t_isolated_gn
|
| 118 |
+
|
| 119 |
+
if cond_channels is not None:
|
| 120 |
+
self.use_spatial_affine = True
|
| 121 |
+
self.norm_layer = DummyAffine(num_channels, affine=affine)
|
| 122 |
+
self.conv_y = SpatialParallelConv3d(
|
| 123 |
+
cond_channels,
|
| 124 |
+
num_channels,
|
| 125 |
+
kernel_size=1,
|
| 126 |
+
padding_mode=padding_mode,
|
| 127 |
+
padding_mode_t=padding_mode_t,
|
| 128 |
+
causal=causal,
|
| 129 |
+
)
|
| 130 |
+
self.conv_b = SpatialParallelConv3d(
|
| 131 |
+
cond_channels,
|
| 132 |
+
num_channels,
|
| 133 |
+
kernel_size=1,
|
| 134 |
+
padding_mode=padding_mode,
|
| 135 |
+
padding_mode_t=padding_mode_t,
|
| 136 |
+
causal=causal,
|
| 137 |
+
)
|
| 138 |
+
else:
|
| 139 |
+
self.use_spatial_affine = False
|
| 140 |
+
if self.affine:
|
| 141 |
+
self.weight = torch.nn.Parameter(torch.ones(num_channels))
|
| 142 |
+
self.bias = torch.nn.Parameter(torch.zeros(num_channels))
|
| 143 |
+
else:
|
| 144 |
+
self.register_parameter("weight", None)
|
| 145 |
+
self.register_parameter("bias", None)
|
| 146 |
+
|
| 147 |
+
def forward(self, f, cond=None):
|
| 148 |
+
need_reshape = self.use_t_isolated_gn and f.dim() == 5
|
| 149 |
+
batch = f.shape[0] if need_reshape else None
|
| 150 |
+
f_size = f.shape[-3:]
|
| 151 |
+
if need_reshape:
|
| 152 |
+
f = _merge_time_to_batch(f)
|
| 153 |
+
|
| 154 |
+
if self.use_spatial_affine:
|
| 155 |
+
scale = self.conv_y(cond)
|
| 156 |
+
bias = self.conv_b(cond)
|
| 157 |
+
if math.prod(scale.shape[-3:]) * math.prod(bias.shape[-3:]) > 1:
|
| 158 |
+
scale = F.interpolate(scale, size=f_size, mode="nearest")
|
| 159 |
+
bias = F.interpolate(bias, size=f_size, mode="nearest")
|
| 160 |
+
if need_reshape:
|
| 161 |
+
scale = _merge_time_to_batch(scale)
|
| 162 |
+
bias = _merge_time_to_batch(bias)
|
| 163 |
+
out = fused_spatial_norm(
|
| 164 |
+
f,
|
| 165 |
+
self.num_groups,
|
| 166 |
+
self.norm_layer.weight,
|
| 167 |
+
self.norm_layer.bias,
|
| 168 |
+
scale,
|
| 169 |
+
bias,
|
| 170 |
+
self.eps,
|
| 171 |
+
self.activation,
|
| 172 |
+
)
|
| 173 |
+
else:
|
| 174 |
+
if cond is not None:
|
| 175 |
+
raise NotImplementedError("Dynamic affine is not defined")
|
| 176 |
+
weight = self.weight if self.affine else None
|
| 177 |
+
bias = self.bias if self.affine else None
|
| 178 |
+
out = fused_group_norm(
|
| 179 |
+
f, self.num_groups, weight, bias, self.eps, self.activation
|
| 180 |
+
)
|
| 181 |
+
|
| 182 |
+
if need_reshape:
|
| 183 |
+
out = _split_time_from_batch(out, batch)
|
| 184 |
+
return out
|
| 185 |
+
|
| 186 |
+
|
| 187 |
+
class SpatialParallelGroupNorm(nn.GroupNorm):
|
| 188 |
+
def __init__(
|
| 189 |
+
self,
|
| 190 |
+
*args,
|
| 191 |
+
**kwargs,
|
| 192 |
+
):
|
| 193 |
+
super().__init__(*args, **kwargs)
|
| 194 |
+
self.spatial_parallel = False
|
| 195 |
+
|
| 196 |
+
def _compute_stats(self, input):
|
| 197 |
+
batch, channels = input.shape[0], input.shape[1]
|
| 198 |
+
spatial_dims = input.shape[2:]
|
| 199 |
+
spatial_size = math.prod(spatial_dims)
|
| 200 |
+
|
| 201 |
+
groups = self.num_groups
|
| 202 |
+
x = input.reshape(batch, groups, channels // groups, -1).to(torch.float32)
|
| 203 |
+
|
| 204 |
+
local_sum = x.sum(dim=(2, 3))
|
| 205 |
+
local_square_sum = (x * x).sum(dim=(2, 3))
|
| 206 |
+
local_n = (channels // groups) * spatial_size
|
| 207 |
+
local_n_tensor = torch.full_like(local_sum, float(local_n))
|
| 208 |
+
|
| 209 |
+
stats = torch.stack([local_sum, local_square_sum, local_n_tensor], dim=0)
|
| 210 |
+
|
| 211 |
+
local_process_group = get_parallel_state()["local_process_group"]
|
| 212 |
+
stats = all_reduce(stats, dist.ReduceOp.SUM, local_process_group)
|
| 213 |
+
|
| 214 |
+
total_sum = stats[0]
|
| 215 |
+
total_square_sum = stats[1]
|
| 216 |
+
total_n = stats[2]
|
| 217 |
+
|
| 218 |
+
mean = total_sum / total_n
|
| 219 |
+
var = (total_square_sum / total_n) - mean**2
|
| 220 |
+
return mean, var
|
| 221 |
+
|
| 222 |
+
def forward(self, input):
|
| 223 |
+
if not self.spatial_parallel:
|
| 224 |
+
return nn.GroupNorm.forward(self, input)
|
| 225 |
+
|
| 226 |
+
batch, channels = input.shape[0], input.shape[1]
|
| 227 |
+
orig_shape = input.shape
|
| 228 |
+
|
| 229 |
+
mean, var = self._compute_stats(input)
|
| 230 |
+
x = input.reshape(batch, self.num_groups, channels // self.num_groups, -1)
|
| 231 |
+
|
| 232 |
+
mean = mean.unsqueeze(-1).unsqueeze(-1)
|
| 233 |
+
var = var.unsqueeze(-1).unsqueeze(-1)
|
| 234 |
+
x = (x - mean) / torch.sqrt(var + self.eps)
|
| 235 |
+
x = x.reshape(orig_shape)
|
| 236 |
+
|
| 237 |
+
if self.affine:
|
| 238 |
+
shape = [1, -1] + [1] * (len(orig_shape) - 2)
|
| 239 |
+
x *= self.weight.view(*shape)
|
| 240 |
+
x += self.bias.view(*shape)
|
| 241 |
+
|
| 242 |
+
return x
|
| 243 |
+
|
| 244 |
+
|
| 245 |
+
class TemporalIsolatedSpatialParallelGroupNorm(SpatialParallelGroupNorm):
|
| 246 |
+
def forward(self, input):
|
| 247 |
+
if input.dim() == 5:
|
| 248 |
+
batch = input.shape[0]
|
| 249 |
+
input = _merge_time_to_batch(input)
|
| 250 |
+
output = super().forward(input)
|
| 251 |
+
return _split_time_from_batch(output, batch)
|
| 252 |
+
return super().forward(input)
|
| 253 |
+
|
| 254 |
+
|
| 255 |
+
|
| 256 |
+
|
| 257 |
+
|
| 258 |
+
|
| 259 |
+
|
| 260 |
+
|
| 261 |
+
class SpatialNorm3D(nn.Module):
|
| 262 |
+
def __init__(
|
| 263 |
+
self,
|
| 264 |
+
f_channels,
|
| 265 |
+
zq_channels,
|
| 266 |
+
padding_mode="zeros",
|
| 267 |
+
padding_mode_t=None,
|
| 268 |
+
causal=True,
|
| 269 |
+
use_t_isolated_gn=False,
|
| 270 |
+
):
|
| 271 |
+
super().__init__()
|
| 272 |
+
norm_cls = (
|
| 273 |
+
TemporalIsolatedSpatialParallelGroupNorm
|
| 274 |
+
if use_t_isolated_gn
|
| 275 |
+
else SpatialParallelGroupNorm
|
| 276 |
+
)
|
| 277 |
+
self.norm_layer = norm_cls(
|
| 278 |
+
num_groups=32, num_channels=f_channels, eps=1e-6, affine=True
|
| 279 |
+
)
|
| 280 |
+
|
| 281 |
+
self.conv_y = SpatialParallelConv3d(
|
| 282 |
+
zq_channels,
|
| 283 |
+
f_channels,
|
| 284 |
+
kernel_size=1,
|
| 285 |
+
padding_mode=padding_mode,
|
| 286 |
+
padding_mode_t=padding_mode_t,
|
| 287 |
+
causal=causal,
|
| 288 |
+
)
|
| 289 |
+
self.conv_b = SpatialParallelConv3d(
|
| 290 |
+
zq_channels,
|
| 291 |
+
f_channels,
|
| 292 |
+
kernel_size=1,
|
| 293 |
+
padding_mode=padding_mode,
|
| 294 |
+
padding_mode_t=padding_mode_t,
|
| 295 |
+
causal=causal,
|
| 296 |
+
)
|
| 297 |
+
|
| 298 |
+
def forward(self, f, zq):
|
| 299 |
+
f_size = f.shape[-3:]
|
| 300 |
+
norm_f = self.norm_layer(f)
|
| 301 |
+
scale = self.conv_y(zq)
|
| 302 |
+
bias = self.conv_b(zq)
|
| 303 |
+
|
| 304 |
+
if math.prod(scale.shape[-3:]) * math.prod(bias.shape[-3:]) > 1:
|
| 305 |
+
scale = F.interpolate(scale, size=f_size, mode="nearest")
|
| 306 |
+
bias = F.interpolate(bias, size=f_size, mode="nearest")
|
| 307 |
+
|
| 308 |
+
return norm_f * scale + bias
|
| 309 |
+
|
| 310 |
+
|
| 311 |
+
def get_spatial_norm_3d(
|
| 312 |
+
num_channels,
|
| 313 |
+
cond_channels,
|
| 314 |
+
*,
|
| 315 |
+
padding_mode="zeros",
|
| 316 |
+
padding_mode_t=None,
|
| 317 |
+
causal=True,
|
| 318 |
+
use_t_isolated_gn=False,
|
| 319 |
+
):
|
| 320 |
+
if os.environ.get("MINIMAX_H3_USE_FUSED_NORM", "false").lower() == "true":
|
| 321 |
+
return FusedGroupNorm3D(
|
| 322 |
+
num_groups=32,
|
| 323 |
+
num_channels=num_channels,
|
| 324 |
+
eps=1e-6,
|
| 325 |
+
affine=True,
|
| 326 |
+
cond_channels=cond_channels,
|
| 327 |
+
use_t_isolated_gn=use_t_isolated_gn,
|
| 328 |
+
padding_mode=padding_mode,
|
| 329 |
+
padding_mode_t=padding_mode_t,
|
| 330 |
+
causal=causal,
|
| 331 |
+
)
|
| 332 |
+
return SpatialNorm3D(
|
| 333 |
+
num_channels,
|
| 334 |
+
cond_channels,
|
| 335 |
+
padding_mode=padding_mode,
|
| 336 |
+
padding_mode_t=padding_mode_t,
|
| 337 |
+
causal=causal,
|
| 338 |
+
use_t_isolated_gn=use_t_isolated_gn,
|
| 339 |
+
)
|
| 340 |
+
|
| 341 |
+
|
| 342 |
+
def get_group_norm_3d(num_channels, use_t_isolated_gn=False):
|
| 343 |
+
if os.environ.get("MINIMAX_H3_USE_FUSED_NORM", "false").lower() == "true":
|
| 344 |
+
return FusedGroupNorm3D(
|
| 345 |
+
num_groups=32,
|
| 346 |
+
num_channels=num_channels,
|
| 347 |
+
eps=1e-6,
|
| 348 |
+
affine=True,
|
| 349 |
+
use_t_isolated_gn=use_t_isolated_gn,
|
| 350 |
+
)
|
| 351 |
+
|
| 352 |
+
norm_cls = (
|
| 353 |
+
TemporalIsolatedSpatialParallelGroupNorm
|
| 354 |
+
if use_t_isolated_gn
|
| 355 |
+
else SpatialParallelGroupNorm
|
| 356 |
+
)
|
| 357 |
+
return norm_cls(num_groups=32, num_channels=num_channels, eps=1e-6, affine=True)
|
FL2VA/video_vae/normalize.py
ADDED
|
@@ -0,0 +1,39 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 2 |
+
# Pixel normalization transforms for the MiniMax H3 visual VAE.
|
| 3 |
+
from typing import Tuple
|
| 4 |
+
from torchvision.transforms import Normalize
|
| 5 |
+
|
| 6 |
+
|
| 7 |
+
NORM_CONFIGS = {
|
| 8 |
+
"imagenet": {
|
| 9 |
+
"mean": (0.485, 0.456, 0.406),
|
| 10 |
+
"std": (0.229, 0.224, 0.225),
|
| 11 |
+
},
|
| 12 |
+
"simple": {
|
| 13 |
+
"mean": (0.5, 0.5, 0.5),
|
| 14 |
+
"std": (0.5, 0.5, 0.5),
|
| 15 |
+
},
|
| 16 |
+
"raw": {
|
| 17 |
+
"mean": (0.0, 0.0, 0.0),
|
| 18 |
+
"std": (1.0, 1.0, 1.0),
|
| 19 |
+
},
|
| 20 |
+
}
|
| 21 |
+
|
| 22 |
+
|
| 23 |
+
def get_norm_constants(norm_type: str = "imagenet") -> Tuple[Tuple[float, ...], Tuple[float, ...]]:
|
| 24 |
+
if norm_type not in NORM_CONFIGS:
|
| 25 |
+
raise ValueError(f"Unknown norm_type: {norm_type}. Must be one of {list(NORM_CONFIGS.keys())}")
|
| 26 |
+
config = NORM_CONFIGS[norm_type]
|
| 27 |
+
return config["mean"], config["std"]
|
| 28 |
+
|
| 29 |
+
|
| 30 |
+
def get_normalize_transform(norm_type: str = "imagenet") -> Normalize:
|
| 31 |
+
mean, std = get_norm_constants(norm_type)
|
| 32 |
+
return Normalize(mean, std)
|
| 33 |
+
|
| 34 |
+
|
| 35 |
+
def get_denormalize_transform(norm_type: str = "imagenet") -> Normalize:
|
| 36 |
+
mean, std = get_norm_constants(norm_type)
|
| 37 |
+
inv_mean = tuple(-m / s for m, s in zip(mean, std))
|
| 38 |
+
inv_std = tuple(1.0 / s for s in std)
|
| 39 |
+
return Normalize(inv_mean, inv_std)
|
FL2VA/video_vae/parallel.py
ADDED
|
@@ -0,0 +1,418 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 2 |
+
# Parallel state and collective helpers for the MiniMax H3 visual VAE.
|
| 3 |
+
import os
|
| 4 |
+
import math
|
| 5 |
+
import torch
|
| 6 |
+
import torch.nn.functional as F
|
| 7 |
+
import torch.distributed as dist
|
| 8 |
+
from torch.autograd import Function
|
| 9 |
+
from torch.distributed import group, ReduceOp
|
| 10 |
+
|
| 11 |
+
|
| 12 |
+
def get_group_rank(group_size):
|
| 13 |
+
global_rank = int(os.environ["RANK"])
|
| 14 |
+
group_rank = global_rank % group_size
|
| 15 |
+
return group_rank
|
| 16 |
+
|
| 17 |
+
|
| 18 |
+
_parallel_state = {}
|
| 19 |
+
|
| 20 |
+
# The torch.autograd.Function subclasses below keep their backward() methods
|
| 21 |
+
# to satisfy the autograd.Function contract; only the forward paths are
|
| 22 |
+
# exercised in this inference-only bundle.
|
| 23 |
+
|
| 24 |
+
|
| 25 |
+
def get_parallel_state():
|
| 26 |
+
return _parallel_state
|
| 27 |
+
|
| 28 |
+
|
| 29 |
+
class _AllGather(Function):
|
| 30 |
+
@staticmethod
|
| 31 |
+
def forward(ctx, group, tensor):
|
| 32 |
+
tensor = tensor.contiguous()
|
| 33 |
+
ctx.group = group
|
| 34 |
+
group_size = dist.get_world_size(group=group)
|
| 35 |
+
out_tensor_list = [torch.empty_like(tensor) for _ in range(group_size)]
|
| 36 |
+
dist.all_gather(out_tensor_list, tensor, group=group)
|
| 37 |
+
return tuple(out_tensor_list)
|
| 38 |
+
|
| 39 |
+
@staticmethod
|
| 40 |
+
def backward(ctx, *grad_outputs):
|
| 41 |
+
rank = dist.get_rank(group=ctx.group)
|
| 42 |
+
gx = torch.empty_like(grad_outputs[rank])
|
| 43 |
+
gx = gx.contiguous()
|
| 44 |
+
grad_outputs = tuple(t.contiguous() for t in grad_outputs)
|
| 45 |
+
dist.reduce_scatter(gx, list(grad_outputs), op=ReduceOp.SUM, group=ctx.group)
|
| 46 |
+
return (None, gx)
|
| 47 |
+
|
| 48 |
+
|
| 49 |
+
@torch.compiler.disable
|
| 50 |
+
def all_gather(tensor, group=group.WORLD):
|
| 51 |
+
return _AllGather.apply(group, tensor)
|
| 52 |
+
|
| 53 |
+
|
| 54 |
+
class _AllGatherVarShape(Function):
|
| 55 |
+
@staticmethod
|
| 56 |
+
def forward(ctx, group, tensor):
|
| 57 |
+
tensor = tensor.contiguous()
|
| 58 |
+
ctx.group = group
|
| 59 |
+
ctx.original_shape = tensor.shape
|
| 60 |
+
|
| 61 |
+
shape_info = torch.tensor(
|
| 62 |
+
list(tensor.shape), dtype=torch.long, device=tensor.device
|
| 63 |
+
)
|
| 64 |
+
|
| 65 |
+
shape_list = [
|
| 66 |
+
torch.empty_like(shape_info)
|
| 67 |
+
for _ in range(dist.get_world_size(group=group))
|
| 68 |
+
]
|
| 69 |
+
dist.all_gather(shape_list, shape_info, group=group)
|
| 70 |
+
|
| 71 |
+
all_shapes = [tuple(shape_tensor.tolist()) for shape_tensor in shape_list]
|
| 72 |
+
ctx.all_shapes = all_shapes
|
| 73 |
+
|
| 74 |
+
flat_tensor = tensor.flatten()
|
| 75 |
+
max_size = max(math.prod(s) for s in all_shapes)
|
| 76 |
+
|
| 77 |
+
if flat_tensor.numel() < max_size:
|
| 78 |
+
padded = torch.zeros(max_size, dtype=tensor.dtype, device=tensor.device)
|
| 79 |
+
padded[: flat_tensor.numel()] = flat_tensor
|
| 80 |
+
flat_tensor = padded
|
| 81 |
+
|
| 82 |
+
gathered_flat = [torch.empty_like(flat_tensor) for _ in range(len(all_shapes))]
|
| 83 |
+
dist.all_gather(gathered_flat, flat_tensor, group=group)
|
| 84 |
+
|
| 85 |
+
return tuple(
|
| 86 |
+
t[: math.prod(shape)].reshape(shape)
|
| 87 |
+
for t, shape in zip(gathered_flat, all_shapes)
|
| 88 |
+
)
|
| 89 |
+
|
| 90 |
+
@staticmethod
|
| 91 |
+
def backward(ctx, *grad_outputs):
|
| 92 |
+
rank = dist.get_rank(group=ctx.group)
|
| 93 |
+
|
| 94 |
+
grad_input = grad_outputs[rank]
|
| 95 |
+
if grad_input is None:
|
| 96 |
+
return None, torch.zeros(
|
| 97 |
+
ctx.original_shape, device=next(iter(grad_outputs)).device
|
| 98 |
+
)
|
| 99 |
+
|
| 100 |
+
max_size = max(math.prod(shape) for shape in ctx.all_shapes)
|
| 101 |
+
padded_grads = []
|
| 102 |
+
|
| 103 |
+
for grad, shape in zip(grad_outputs, ctx.all_shapes):
|
| 104 |
+
if grad is not None:
|
| 105 |
+
flat_grad = grad.flatten()
|
| 106 |
+
else:
|
| 107 |
+
flat_grad = torch.zeros(
|
| 108 |
+
math.prod(shape),
|
| 109 |
+
dtype=grad_input.dtype,
|
| 110 |
+
device=grad_input.device,
|
| 111 |
+
)
|
| 112 |
+
|
| 113 |
+
if flat_grad.numel() < max_size:
|
| 114 |
+
padded = torch.zeros(
|
| 115 |
+
max_size, dtype=flat_grad.dtype, device=flat_grad.device
|
| 116 |
+
)
|
| 117 |
+
padded[: flat_grad.numel()] = flat_grad
|
| 118 |
+
padded_grads.append(padded)
|
| 119 |
+
else:
|
| 120 |
+
padded_grads.append(flat_grad)
|
| 121 |
+
|
| 122 |
+
result_grad = torch.empty_like(padded_grads[0])
|
| 123 |
+
dist.reduce_scatter(result_grad, padded_grads, op=ReduceOp.SUM, group=ctx.group)
|
| 124 |
+
|
| 125 |
+
original_size = math.prod(ctx.original_shape)
|
| 126 |
+
return None, result_grad[:original_size].reshape(ctx.original_shape)
|
| 127 |
+
|
| 128 |
+
|
| 129 |
+
@torch.compiler.disable
|
| 130 |
+
def all_gather_var_shape(tensor, group=group.WORLD):
|
| 131 |
+
return _AllGatherVarShape.apply(group, tensor)
|
| 132 |
+
|
| 133 |
+
|
| 134 |
+
class _AllReduce(Function):
|
| 135 |
+
@staticmethod
|
| 136 |
+
def forward(ctx, _input, op, group):
|
| 137 |
+
ctx.group = group
|
| 138 |
+
ctx.op = op
|
| 139 |
+
_input = _input.clone()
|
| 140 |
+
dist.all_reduce(_input, op=op, group=group)
|
| 141 |
+
return _input
|
| 142 |
+
|
| 143 |
+
@staticmethod
|
| 144 |
+
def backward(ctx, grad_output):
|
| 145 |
+
grad_output = grad_output.clone()
|
| 146 |
+
dist.all_reduce(grad_output, op=ctx.op, group=ctx.group)
|
| 147 |
+
return grad_output, None, None
|
| 148 |
+
|
| 149 |
+
|
| 150 |
+
@torch.compiler.disable
|
| 151 |
+
def all_reduce(input_, op, group):
|
| 152 |
+
return _AllReduce.apply(input_, op, group)
|
| 153 |
+
|
| 154 |
+
|
| 155 |
+
class _AlltoAllSingle(Function):
|
| 156 |
+
@staticmethod
|
| 157 |
+
def forward(ctx, group, input):
|
| 158 |
+
ctx.group = group
|
| 159 |
+
|
| 160 |
+
world_size = dist.get_world_size(group=group)
|
| 161 |
+
if world_size == 1:
|
| 162 |
+
return input
|
| 163 |
+
|
| 164 |
+
input = input.contiguous()
|
| 165 |
+
output = torch.empty_like(input)
|
| 166 |
+
dist.all_to_all_single(
|
| 167 |
+
output,
|
| 168 |
+
input,
|
| 169 |
+
group=group,
|
| 170 |
+
)
|
| 171 |
+
return output
|
| 172 |
+
|
| 173 |
+
@staticmethod
|
| 174 |
+
def backward(ctx, grad_output):
|
| 175 |
+
return (None, _AlltoAllSingle.apply(ctx.group, grad_output))
|
| 176 |
+
|
| 177 |
+
|
| 178 |
+
@torch.compiler.disable
|
| 179 |
+
def all_to_all_single(input, group=group.WORLD):
|
| 180 |
+
return _AlltoAllSingle.apply(group, input)
|
| 181 |
+
|
| 182 |
+
|
| 183 |
+
|
| 184 |
+
@torch.compiler.disable
|
| 185 |
+
def get_subseq(input, sp_size=None):
|
| 186 |
+
if sp_size is None:
|
| 187 |
+
state = get_parallel_state()
|
| 188 |
+
if not state.get("sp_enabled", False):
|
| 189 |
+
return input
|
| 190 |
+
sp_size = state["sp_size"]
|
| 191 |
+
sp_rank = state["sp_rank"]
|
| 192 |
+
else:
|
| 193 |
+
sp_rank = get_group_rank(sp_size)
|
| 194 |
+
|
| 195 |
+
if sp_size == 1:
|
| 196 |
+
return input
|
| 197 |
+
|
| 198 |
+
if input.shape[1] % sp_size != 0:
|
| 199 |
+
raise ValueError(
|
| 200 |
+
f"Input shape {input.shape} is not divisible by sp_size {sp_size}"
|
| 201 |
+
)
|
| 202 |
+
|
| 203 |
+
return torch.chunk(input, sp_size, dim=1)[sp_rank]
|
| 204 |
+
|
| 205 |
+
|
| 206 |
+
@torch.compiler.disable
|
| 207 |
+
def gather_subseq(input, sp_size=None, local_process_group=None):
|
| 208 |
+
if sp_size is None:
|
| 209 |
+
state = get_parallel_state()
|
| 210 |
+
if not state.get("sp_enabled", False):
|
| 211 |
+
return input
|
| 212 |
+
sp_size = state["sp_size"]
|
| 213 |
+
local_process_group = state["sp_process_group"]
|
| 214 |
+
|
| 215 |
+
if sp_size == 1:
|
| 216 |
+
return input
|
| 217 |
+
|
| 218 |
+
output = all_gather(input, group=local_process_group)
|
| 219 |
+
output = torch.cat(output, dim=1)
|
| 220 |
+
return output
|
| 221 |
+
|
| 222 |
+
|
| 223 |
+
@torch.compiler.disable
|
| 224 |
+
def all_to_all_4D(
|
| 225 |
+
input: torch.tensor,
|
| 226 |
+
scatter_idx: int = 2,
|
| 227 |
+
gather_idx: int = 1,
|
| 228 |
+
group=None,
|
| 229 |
+
):
|
| 230 |
+
assert (
|
| 231 |
+
input.dim() == 4
|
| 232 |
+
), f"input must be 4D tensor, got {input.dim()} and shape {input.shape}"
|
| 233 |
+
|
| 234 |
+
if group is None:
|
| 235 |
+
seq_world_size = 1
|
| 236 |
+
else:
|
| 237 |
+
seq_world_size = dist.get_world_size(group)
|
| 238 |
+
|
| 239 |
+
if seq_world_size == 1:
|
| 240 |
+
return input
|
| 241 |
+
|
| 242 |
+
if scatter_idx == 2 and gather_idx == 1:
|
| 243 |
+
bs, shard_seqlen, hc, hs = input.shape
|
| 244 |
+
seqlen = shard_seqlen * seq_world_size
|
| 245 |
+
shard_hc = hc // seq_world_size
|
| 246 |
+
|
| 247 |
+
input_t = (
|
| 248 |
+
input.reshape(bs, shard_seqlen, seq_world_size, shard_hc, hs)
|
| 249 |
+
.transpose(0, 2)
|
| 250 |
+
.contiguous()
|
| 251 |
+
)
|
| 252 |
+
|
| 253 |
+
output = all_to_all_single(input_t, group=group)
|
| 254 |
+
output = output.reshape(seqlen, bs, shard_hc, hs)
|
| 255 |
+
output = output.transpose(0, 1).contiguous().reshape(bs, seqlen, shard_hc, hs)
|
| 256 |
+
return output
|
| 257 |
+
|
| 258 |
+
elif scatter_idx == 1 and gather_idx == 2:
|
| 259 |
+
bs, seqlen, shard_hc, hs = input.shape
|
| 260 |
+
hc = shard_hc * seq_world_size
|
| 261 |
+
shard_seqlen = seqlen // seq_world_size
|
| 262 |
+
|
| 263 |
+
input_t = (
|
| 264 |
+
input.reshape(bs, seq_world_size, shard_seqlen, shard_hc, hs)
|
| 265 |
+
.transpose(0, 3)
|
| 266 |
+
.transpose(0, 1)
|
| 267 |
+
.contiguous()
|
| 268 |
+
.reshape(seq_world_size, shard_hc, shard_seqlen, bs, hs)
|
| 269 |
+
)
|
| 270 |
+
|
| 271 |
+
output = all_to_all_single(input_t, group=group)
|
| 272 |
+
output = output.reshape(hc, shard_seqlen, bs, hs)
|
| 273 |
+
output = output.transpose(0, 2).contiguous().reshape(bs, shard_seqlen, hc, hs)
|
| 274 |
+
return output
|
| 275 |
+
else:
|
| 276 |
+
raise RuntimeError("scatter_idx must be 1 or 2 and gather_idx must be 1 or 2")
|
| 277 |
+
|
| 278 |
+
|
| 279 |
+
|
| 280 |
+
@torch.compiler.disable
|
| 281 |
+
def exchange_borders(
|
| 282 |
+
input_, padding, pad_mode, sp_rank, sp_size, group, dim=-1, async_op=False
|
| 283 |
+
):
|
| 284 |
+
if async_op and input_.requires_grad:
|
| 285 |
+
raise ValueError("async_op is not supported backward, check previous commits")
|
| 286 |
+
|
| 287 |
+
slice_indices = [slice(None)] * input_.ndim
|
| 288 |
+
slice_indices[dim] = slice(None, padding)
|
| 289 |
+
first_tensor = input_[tuple(slice_indices)].contiguous()
|
| 290 |
+
|
| 291 |
+
slice_indices[dim] = slice(-padding, None)
|
| 292 |
+
last_tensor = input_[tuple(slice_indices)].contiguous()
|
| 293 |
+
|
| 294 |
+
if async_op:
|
| 295 |
+
first_borders = [torch.empty_like(first_tensor) for _ in range(sp_size)]
|
| 296 |
+
last_borders = [torch.empty_like(last_tensor) for _ in range(sp_size)]
|
| 297 |
+
|
| 298 |
+
handle_first = dist.all_gather(
|
| 299 |
+
first_borders, first_tensor, group=group, async_op=True
|
| 300 |
+
)
|
| 301 |
+
handle_last = dist.all_gather(
|
| 302 |
+
last_borders, last_tensor, group=group, async_op=True
|
| 303 |
+
)
|
| 304 |
+
else:
|
| 305 |
+
first_borders = all_gather(first_tensor, group=group)
|
| 306 |
+
last_borders = all_gather(last_tensor, group=group)
|
| 307 |
+
|
| 308 |
+
if dim < 0:
|
| 309 |
+
pad_dim = -1 - dim
|
| 310 |
+
else:
|
| 311 |
+
pad_dim = input_.ndim - 1 - dim
|
| 312 |
+
|
| 313 |
+
pad_size = [0] * ((input_.ndim - 2) * 2)
|
| 314 |
+
pad_size[pad_dim * 2] = padding
|
| 315 |
+
pad_size[pad_dim * 2 + 1] = padding
|
| 316 |
+
output = F.pad(input_, pad_size, mode=pad_mode)
|
| 317 |
+
|
| 318 |
+
slice_indices = [slice(None)] * input_.ndim
|
| 319 |
+
slice_indices[dim] = slice(-padding, None)
|
| 320 |
+
|
| 321 |
+
if async_op:
|
| 322 |
+
handle_first.wait()
|
| 323 |
+
|
| 324 |
+
if sp_rank < sp_size - 1:
|
| 325 |
+
output[tuple(slice_indices)] = first_borders[sp_rank + 1]
|
| 326 |
+
else:
|
| 327 |
+
output[tuple(slice_indices)] += first_borders[0] * 0.0
|
| 328 |
+
|
| 329 |
+
slice_indices = [slice(None)] * input_.ndim
|
| 330 |
+
slice_indices[dim] = slice(None, padding)
|
| 331 |
+
|
| 332 |
+
if async_op:
|
| 333 |
+
handle_last.wait()
|
| 334 |
+
|
| 335 |
+
if sp_rank > 0:
|
| 336 |
+
output[tuple(slice_indices)] = last_borders[sp_rank - 1]
|
| 337 |
+
else:
|
| 338 |
+
output[tuple(slice_indices)] += last_borders[sp_size - 1] * 0.0
|
| 339 |
+
|
| 340 |
+
return output
|
| 341 |
+
|
| 342 |
+
|
| 343 |
+
@torch.compiler.disable
|
| 344 |
+
def exchange_strides(
|
| 345 |
+
input_, pad_mode, sp_rank, sp_size, group, dim=-1, async_op=False
|
| 346 |
+
):
|
| 347 |
+
if async_op and input_.requires_grad:
|
| 348 |
+
raise ValueError("async_op is not supported backward, check previous commits")
|
| 349 |
+
|
| 350 |
+
if dim not in [-1, -2]:
|
| 351 |
+
raise ValueError("dim must be -1 (W) or -2 (H) for exchange_strides")
|
| 352 |
+
|
| 353 |
+
if dim == -1:
|
| 354 |
+
if input_.ndim == 5:
|
| 355 |
+
input_ = F.pad(input_, (0, 0, 0, 1, 0, 0), mode=pad_mode)
|
| 356 |
+
elif input_.ndim == 4:
|
| 357 |
+
input_ = F.pad(input_, (0, 0, 0, 1), mode=pad_mode)
|
| 358 |
+
else:
|
| 359 |
+
raise ValueError(f"Input must have 4 or 5 dimensions, got {input_.ndim}")
|
| 360 |
+
|
| 361 |
+
left_border = input_[..., :1].contiguous()
|
| 362 |
+
|
| 363 |
+
if async_op:
|
| 364 |
+
left_borders = [torch.empty_like(left_border) for _ in range(sp_size)]
|
| 365 |
+
handle = dist.all_gather(
|
| 366 |
+
left_borders, left_border, group=group, async_op=True
|
| 367 |
+
)
|
| 368 |
+
else:
|
| 369 |
+
left_borders = all_gather(left_border, group=group)
|
| 370 |
+
|
| 371 |
+
if input_.ndim == 5:
|
| 372 |
+
output = F.pad(input_, (0, 1, 0, 0, 0, 0), mode=pad_mode)
|
| 373 |
+
elif input_.ndim == 4:
|
| 374 |
+
output = F.pad(input_, (0, 1, 0, 0), mode=pad_mode)
|
| 375 |
+
else:
|
| 376 |
+
raise ValueError(f"Input must have 4 or 5 dimensions, got {input_.ndim}")
|
| 377 |
+
|
| 378 |
+
if async_op:
|
| 379 |
+
handle.wait()
|
| 380 |
+
|
| 381 |
+
if sp_rank != sp_size - 1:
|
| 382 |
+
output[..., -1:] = left_borders[sp_rank + 1]
|
| 383 |
+
else:
|
| 384 |
+
output[..., -1:] += left_borders[0] * 0.0
|
| 385 |
+
else:
|
| 386 |
+
if input_.ndim == 5:
|
| 387 |
+
input_ = F.pad(input_, (0, 1, 0, 0, 0, 0), mode=pad_mode)
|
| 388 |
+
elif input_.ndim == 4:
|
| 389 |
+
input_ = F.pad(input_, (0, 1, 0, 0), mode=pad_mode)
|
| 390 |
+
else:
|
| 391 |
+
raise ValueError(f"Input must have 4 or 5 dimensions, got {input_.ndim}")
|
| 392 |
+
|
| 393 |
+
top_border = input_[..., :1, :].contiguous()
|
| 394 |
+
|
| 395 |
+
if async_op:
|
| 396 |
+
top_borders = [torch.empty_like(top_border) for _ in range(sp_size)]
|
| 397 |
+
handle = dist.all_gather(
|
| 398 |
+
top_borders, top_border, group=group, async_op=True
|
| 399 |
+
)
|
| 400 |
+
else:
|
| 401 |
+
top_borders = all_gather(top_border, group=group)
|
| 402 |
+
|
| 403 |
+
if input_.ndim == 5:
|
| 404 |
+
output = F.pad(input_, (0, 0, 0, 1, 0, 0), mode=pad_mode)
|
| 405 |
+
elif input_.ndim == 4:
|
| 406 |
+
output = F.pad(input_, (0, 0, 0, 1), mode=pad_mode)
|
| 407 |
+
else:
|
| 408 |
+
raise ValueError(f"Input must have 4 or 5 dimensions, got {input_.ndim}")
|
| 409 |
+
|
| 410 |
+
if async_op:
|
| 411 |
+
handle.wait()
|
| 412 |
+
|
| 413 |
+
if sp_rank != sp_size - 1:
|
| 414 |
+
output[..., -1:, :] = top_borders[sp_rank + 1]
|
| 415 |
+
else:
|
| 416 |
+
output[..., -1:, :] += top_borders[0] * 0.0
|
| 417 |
+
|
| 418 |
+
return output
|
FL2VA/video_vae/source/config.json
ADDED
|
@@ -0,0 +1,71 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"_class_name": "AutoencoderKLLegacy",
|
| 3 |
+
"_diffusers_version": "0.32.2",
|
| 4 |
+
"causal_decoder": false,
|
| 5 |
+
"causal_encoder": true,
|
| 6 |
+
"ch": 128,
|
| 7 |
+
"ch_mult": [
|
| 8 |
+
1,
|
| 9 |
+
2,
|
| 10 |
+
2,
|
| 11 |
+
4,
|
| 12 |
+
4,
|
| 13 |
+
8
|
| 14 |
+
],
|
| 15 |
+
"embed_dim": 24,
|
| 16 |
+
"in_channels": 3,
|
| 17 |
+
"num_res_blocks": 2,
|
| 18 |
+
"num_res_blocks_decoder": null,
|
| 19 |
+
"out_ch": 3,
|
| 20 |
+
"padding_mode": "reflect",
|
| 21 |
+
"padding_mode_t": null,
|
| 22 |
+
"pixel_norm_type": "imagenet",
|
| 23 |
+
"scaling_factor": 1.0,
|
| 24 |
+
"shift_factor": 0.0,
|
| 25 |
+
"space_down": [
|
| 26 |
+
2,
|
| 27 |
+
2,
|
| 28 |
+
2,
|
| 29 |
+
2,
|
| 30 |
+
1,
|
| 31 |
+
1
|
| 32 |
+
],
|
| 33 |
+
"space_up": [
|
| 34 |
+
1,
|
| 35 |
+
2,
|
| 36 |
+
2,
|
| 37 |
+
2,
|
| 38 |
+
2,
|
| 39 |
+
1
|
| 40 |
+
],
|
| 41 |
+
"time_down": [
|
| 42 |
+
1,
|
| 43 |
+
2,
|
| 44 |
+
2,
|
| 45 |
+
1,
|
| 46 |
+
1,
|
| 47 |
+
1
|
| 48 |
+
],
|
| 49 |
+
"time_up": null,
|
| 50 |
+
"use_3d_conv": true,
|
| 51 |
+
"use_t_isolated_gn": true,
|
| 52 |
+
"use_vit_decoder": true,
|
| 53 |
+
"vae_ratio": 16,
|
| 54 |
+
"vae_ratio_t": 4,
|
| 55 |
+
"vit_decoder_kwargs": {
|
| 56 |
+
"dim_head": 64,
|
| 57 |
+
"ffn_activation_fn": "silu",
|
| 58 |
+
"ffn_use_gated": true,
|
| 59 |
+
"heads": 32,
|
| 60 |
+
"norm_affine": true,
|
| 61 |
+
"norm_type": "rms_norm",
|
| 62 |
+
"num_layers": 36,
|
| 63 |
+
"qk_norm_affine": false,
|
| 64 |
+
"qk_norm_type": "rms_norm",
|
| 65 |
+
"rope_dim_ratio": 0.75,
|
| 66 |
+
"rope_theta": 100.0
|
| 67 |
+
},
|
| 68 |
+
"z_channels": 24,
|
| 69 |
+
"zq_ch_decoder": null,
|
| 70 |
+
"zq_ch_encoder": null
|
| 71 |
+
}
|
FL2VA/video_vae/source/model.safetensors
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:5f0c2e161d895a9fee7645ca32d4a7e3a22b90cacfcbeba62ec999cdbbefe0d3
|
| 3 |
+
size 10415548320
|
FL2VA/video_vae/utils.py
ADDED
|
@@ -0,0 +1,18 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 2 |
+
# Module helpers for the MiniMax H3 visual VAE (inference-only bundle).
|
| 3 |
+
|
| 4 |
+
|
| 5 |
+
def apply_spatial_parallel(module, enabled, chunk_dim=-1):
|
| 6 |
+
from .conv import SpatialParallelConv3d
|
| 7 |
+
from .norm import FusedGroupNorm3D, SpatialParallelGroupNorm
|
| 8 |
+
|
| 9 |
+
if hasattr(module, "set_spatial_parallel"):
|
| 10 |
+
module.set_spatial_parallel(enabled)
|
| 11 |
+
for m in module.modules():
|
| 12 |
+
if enabled and isinstance(m, FusedGroupNorm3D):
|
| 13 |
+
raise NotImplementedError("FusedGroupNorm3D is incompatible with SP")
|
| 14 |
+
if isinstance(m, SpatialParallelGroupNorm):
|
| 15 |
+
m.spatial_parallel = enabled
|
| 16 |
+
elif isinstance(m, SpatialParallelConv3d):
|
| 17 |
+
m.spatial_parallel = enabled
|
| 18 |
+
m.chunk_dim = chunk_dim
|
FL2VA/video_vae/vae_cnn.py
ADDED
|
@@ -0,0 +1,304 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 2 |
+
# 3D causal CNN encoder for the MiniMax H3 visual VAE (inference-only bundle).
|
| 3 |
+
import os
|
| 4 |
+
import torch.nn as nn
|
| 5 |
+
import torch.nn.functional as F
|
| 6 |
+
|
| 7 |
+
from .attention import maybe_checkpoint
|
| 8 |
+
from .conv import SpatialParallelConv3d
|
| 9 |
+
from .norm import get_spatial_norm_3d
|
| 10 |
+
from .parallel import get_parallel_state, exchange_strides
|
| 11 |
+
from .norm import get_group_norm_3d
|
| 12 |
+
|
| 13 |
+
|
| 14 |
+
|
| 15 |
+
|
| 16 |
+
|
| 17 |
+
|
| 18 |
+
|
| 19 |
+
|
| 20 |
+
|
| 21 |
+
|
| 22 |
+
# ============================================================================
|
| 23 |
+
# 3D CNN Components
|
| 24 |
+
# ============================================================================
|
| 25 |
+
|
| 26 |
+
|
| 27 |
+
def norm_silu(x, norm, cond=None):
|
| 28 |
+
if cond is None:
|
| 29 |
+
return F.silu(norm(x))
|
| 30 |
+
else:
|
| 31 |
+
return F.silu(norm(x, cond))
|
| 32 |
+
|
| 33 |
+
|
| 34 |
+
class Downsample3D(nn.Module):
|
| 35 |
+
def __init__(
|
| 36 |
+
self,
|
| 37 |
+
in_channels,
|
| 38 |
+
out_channels,
|
| 39 |
+
time_stride=1,
|
| 40 |
+
space_stride=2,
|
| 41 |
+
padding_mode="zeros",
|
| 42 |
+
padding_mode_t=None,
|
| 43 |
+
causal=True,
|
| 44 |
+
):
|
| 45 |
+
super().__init__()
|
| 46 |
+
self.time_stride = time_stride
|
| 47 |
+
self.space_stride = space_stride
|
| 48 |
+
|
| 49 |
+
assert time_stride in [1, 2]
|
| 50 |
+
assert space_stride in [1, 2, 3]
|
| 51 |
+
|
| 52 |
+
self.conv = SpatialParallelConv3d(
|
| 53 |
+
in_channels,
|
| 54 |
+
out_channels,
|
| 55 |
+
kernel_size=3,
|
| 56 |
+
padding=(1, 0, 0),
|
| 57 |
+
stride=(time_stride, space_stride, space_stride),
|
| 58 |
+
padding_mode=padding_mode,
|
| 59 |
+
padding_mode_t=padding_mode_t,
|
| 60 |
+
causal=causal,
|
| 61 |
+
)
|
| 62 |
+
self.causal = self.conv.causal
|
| 63 |
+
self.pad_mode = self.conv.pad_mode
|
| 64 |
+
|
| 65 |
+
def forward(self, x):
|
| 66 |
+
if self.space_stride == 2:
|
| 67 |
+
if getattr(self.conv, "spatial_parallel", False):
|
| 68 |
+
state = get_parallel_state()
|
| 69 |
+
x = exchange_strides(
|
| 70 |
+
x,
|
| 71 |
+
self.pad_mode,
|
| 72 |
+
state["sp_rank"],
|
| 73 |
+
state["sp_size"],
|
| 74 |
+
state["sp_process_group"],
|
| 75 |
+
self.conv.chunk_dim,
|
| 76 |
+
)
|
| 77 |
+
else:
|
| 78 |
+
pad = (0, 1, 0, 1, 0, 0)
|
| 79 |
+
x = F.pad(x, pad, mode=self.pad_mode)
|
| 80 |
+
return self.conv(x)
|
| 81 |
+
|
| 82 |
+
|
| 83 |
+
class ResnetBlock3D(nn.Module):
|
| 84 |
+
def __init__(
|
| 85 |
+
self,
|
| 86 |
+
in_channels,
|
| 87 |
+
out_channels=None,
|
| 88 |
+
zq_ch=None,
|
| 89 |
+
padding_mode="zeros",
|
| 90 |
+
padding_mode_t=None,
|
| 91 |
+
causal=True,
|
| 92 |
+
use_t_isolated_gn=False,
|
| 93 |
+
):
|
| 94 |
+
super().__init__()
|
| 95 |
+
self.in_channels = in_channels
|
| 96 |
+
out_channels = in_channels if out_channels is None else out_channels
|
| 97 |
+
self.out_channels = out_channels
|
| 98 |
+
|
| 99 |
+
self.use_fused_norm = (
|
| 100 |
+
os.environ.get("MINIMAX_H3_USE_FUSED_NORM", "false").lower() == "true"
|
| 101 |
+
)
|
| 102 |
+
|
| 103 |
+
if zq_ch is None:
|
| 104 |
+
self.norm1 = get_group_norm_3d(in_channels, use_t_isolated_gn=use_t_isolated_gn)
|
| 105 |
+
self.norm2 = get_group_norm_3d(out_channels, use_t_isolated_gn=use_t_isolated_gn)
|
| 106 |
+
else:
|
| 107 |
+
self.norm1 = get_spatial_norm_3d(
|
| 108 |
+
in_channels,
|
| 109 |
+
zq_ch,
|
| 110 |
+
padding_mode=padding_mode,
|
| 111 |
+
padding_mode_t=padding_mode_t,
|
| 112 |
+
causal=causal,
|
| 113 |
+
use_t_isolated_gn=use_t_isolated_gn,
|
| 114 |
+
)
|
| 115 |
+
self.norm2 = get_spatial_norm_3d(
|
| 116 |
+
out_channels,
|
| 117 |
+
zq_ch,
|
| 118 |
+
padding_mode=padding_mode,
|
| 119 |
+
padding_mode_t=padding_mode_t,
|
| 120 |
+
causal=causal,
|
| 121 |
+
use_t_isolated_gn=use_t_isolated_gn,
|
| 122 |
+
)
|
| 123 |
+
|
| 124 |
+
self.conv1 = SpatialParallelConv3d(
|
| 125 |
+
in_channels,
|
| 126 |
+
out_channels,
|
| 127 |
+
kernel_size=3,
|
| 128 |
+
padding=1,
|
| 129 |
+
padding_mode=padding_mode,
|
| 130 |
+
padding_mode_t=padding_mode_t,
|
| 131 |
+
causal=causal,
|
| 132 |
+
)
|
| 133 |
+
|
| 134 |
+
self.conv2 = SpatialParallelConv3d(
|
| 135 |
+
out_channels,
|
| 136 |
+
out_channels,
|
| 137 |
+
kernel_size=3,
|
| 138 |
+
padding=1,
|
| 139 |
+
padding_mode=padding_mode,
|
| 140 |
+
padding_mode_t=padding_mode_t,
|
| 141 |
+
causal=causal,
|
| 142 |
+
)
|
| 143 |
+
|
| 144 |
+
if self.in_channels != self.out_channels:
|
| 145 |
+
self.nin_shortcut = SpatialParallelConv3d(
|
| 146 |
+
in_channels,
|
| 147 |
+
out_channels,
|
| 148 |
+
kernel_size=1,
|
| 149 |
+
padding_mode=padding_mode,
|
| 150 |
+
padding_mode_t=padding_mode_t,
|
| 151 |
+
causal=causal,
|
| 152 |
+
)
|
| 153 |
+
|
| 154 |
+
def forward(self, x, zq=None):
|
| 155 |
+
h = x
|
| 156 |
+
|
| 157 |
+
if self.use_fused_norm:
|
| 158 |
+
h = self.norm1(h, zq)
|
| 159 |
+
else:
|
| 160 |
+
h = norm_silu(h, self.norm1, zq)
|
| 161 |
+
|
| 162 |
+
h = self.conv1(h)
|
| 163 |
+
|
| 164 |
+
if self.use_fused_norm:
|
| 165 |
+
h = self.norm2(h, zq)
|
| 166 |
+
else:
|
| 167 |
+
h = norm_silu(h, self.norm2, zq)
|
| 168 |
+
|
| 169 |
+
h = self.conv2(h)
|
| 170 |
+
|
| 171 |
+
if self.in_channels != self.out_channels:
|
| 172 |
+
x = self.nin_shortcut(x)
|
| 173 |
+
|
| 174 |
+
return x + h
|
| 175 |
+
|
| 176 |
+
|
| 177 |
+
class EncoderFCN3D(nn.Module):
|
| 178 |
+
def __init__(
|
| 179 |
+
self,
|
| 180 |
+
ch,
|
| 181 |
+
ch_mult,
|
| 182 |
+
space_down,
|
| 183 |
+
time_down,
|
| 184 |
+
num_res_blocks,
|
| 185 |
+
in_channels,
|
| 186 |
+
z_channels,
|
| 187 |
+
double_z=False,
|
| 188 |
+
zq_ch=None,
|
| 189 |
+
padding_mode="zeros",
|
| 190 |
+
padding_mode_t=None,
|
| 191 |
+
causal=True,
|
| 192 |
+
use_t_isolated_gn=False,
|
| 193 |
+
):
|
| 194 |
+
super().__init__()
|
| 195 |
+
self.ch = ch
|
| 196 |
+
self.num_levels = len(ch_mult)
|
| 197 |
+
|
| 198 |
+
if isinstance(num_res_blocks, int):
|
| 199 |
+
self.num_res_blocks = [num_res_blocks] * self.num_levels
|
| 200 |
+
else:
|
| 201 |
+
self.num_res_blocks = num_res_blocks
|
| 202 |
+
|
| 203 |
+
self.space_down_factors = space_down
|
| 204 |
+
self.time_down_factors = time_down
|
| 205 |
+
self.in_channels = in_channels
|
| 206 |
+
|
| 207 |
+
self.use_fused_norm = (
|
| 208 |
+
os.environ.get("MINIMAX_H3_USE_FUSED_NORM", "false").lower() == "true"
|
| 209 |
+
)
|
| 210 |
+
|
| 211 |
+
block_mid = [ch * ch_mult[i] for i in range(self.num_levels)]
|
| 212 |
+
block_in = [block_mid[0]] + block_mid[:-1]
|
| 213 |
+
block_out = block_mid
|
| 214 |
+
|
| 215 |
+
conv_kwargs = dict(
|
| 216 |
+
padding_mode=padding_mode,
|
| 217 |
+
padding_mode_t=padding_mode_t,
|
| 218 |
+
causal=causal,
|
| 219 |
+
)
|
| 220 |
+
|
| 221 |
+
self.conv_in = SpatialParallelConv3d(
|
| 222 |
+
in_channels, block_in[0], kernel_size=3, padding=1, **conv_kwargs
|
| 223 |
+
)
|
| 224 |
+
|
| 225 |
+
self.down = nn.ModuleList()
|
| 226 |
+
for i_level in range(self.num_levels):
|
| 227 |
+
down = nn.Module()
|
| 228 |
+
|
| 229 |
+
down.block = nn.ModuleList()
|
| 230 |
+
for i in range(self.num_res_blocks[i_level]):
|
| 231 |
+
down.block.append(
|
| 232 |
+
ResnetBlock3D(
|
| 233 |
+
in_channels=block_in[i_level] if i == 0 else block_mid[i_level],
|
| 234 |
+
out_channels=block_mid[i_level],
|
| 235 |
+
zq_ch=zq_ch,
|
| 236 |
+
use_t_isolated_gn=use_t_isolated_gn,
|
| 237 |
+
**conv_kwargs,
|
| 238 |
+
)
|
| 239 |
+
)
|
| 240 |
+
|
| 241 |
+
if space_down[i_level] * time_down[i_level] > 1:
|
| 242 |
+
down.downsample = Downsample3D(
|
| 243 |
+
block_mid[i_level],
|
| 244 |
+
block_out[i_level],
|
| 245 |
+
time_stride=time_down[i_level],
|
| 246 |
+
space_stride=space_down[i_level],
|
| 247 |
+
**conv_kwargs,
|
| 248 |
+
)
|
| 249 |
+
else:
|
| 250 |
+
if block_out[i_level] != block_mid[i_level]:
|
| 251 |
+
down.downsample = SpatialParallelConv3d(
|
| 252 |
+
block_mid[i_level],
|
| 253 |
+
block_out[i_level],
|
| 254 |
+
kernel_size=1,
|
| 255 |
+
**conv_kwargs,
|
| 256 |
+
)
|
| 257 |
+
|
| 258 |
+
self.down.append(down)
|
| 259 |
+
|
| 260 |
+
if zq_ch is None:
|
| 261 |
+
self.norm_out = get_group_norm_3d(
|
| 262 |
+
block_out[-1], use_t_isolated_gn=use_t_isolated_gn
|
| 263 |
+
)
|
| 264 |
+
else:
|
| 265 |
+
self.norm_out = get_spatial_norm_3d(
|
| 266 |
+
block_out[-1],
|
| 267 |
+
zq_ch,
|
| 268 |
+
use_t_isolated_gn=use_t_isolated_gn,
|
| 269 |
+
**conv_kwargs,
|
| 270 |
+
)
|
| 271 |
+
|
| 272 |
+
self.conv_out = SpatialParallelConv3d(
|
| 273 |
+
block_out[-1],
|
| 274 |
+
2 * z_channels if double_z else z_channels,
|
| 275 |
+
kernel_size=3,
|
| 276 |
+
padding=1,
|
| 277 |
+
**conv_kwargs,
|
| 278 |
+
)
|
| 279 |
+
|
| 280 |
+
self.gradient_checkpointing = False
|
| 281 |
+
|
| 282 |
+
def _set_gradient_checkpointing(self, module, value=False):
|
| 283 |
+
if hasattr(module, "gradient_checkpointing"):
|
| 284 |
+
module.gradient_checkpointing = value
|
| 285 |
+
|
| 286 |
+
def forward(self, x, zq=None):
|
| 287 |
+
h = self.conv_in(x)
|
| 288 |
+
for i_level in range(self.num_levels):
|
| 289 |
+
for i_block in range(self.num_res_blocks[i_level]):
|
| 290 |
+
h = maybe_checkpoint(self, self.down[i_level].block[i_block], h, zq)
|
| 291 |
+
if hasattr(self.down[i_level], "downsample"):
|
| 292 |
+
h = self.down[i_level].downsample(h)
|
| 293 |
+
|
| 294 |
+
if self.use_fused_norm:
|
| 295 |
+
h = self.norm_out(h, zq)
|
| 296 |
+
else:
|
| 297 |
+
h = norm_silu(h, self.norm_out, zq)
|
| 298 |
+
|
| 299 |
+
h = self.conv_out(h)
|
| 300 |
+
return h
|
| 301 |
+
|
| 302 |
+
|
| 303 |
+
|
| 304 |
+
|
FL2VA/video_vae/vae_module.py
ADDED
|
@@ -0,0 +1,53 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 2 |
+
# VAE distribution and aggregation helpers for the MiniMax H3 visual VAE.
|
| 3 |
+
import torch
|
| 4 |
+
|
| 5 |
+
|
| 6 |
+
class DiagonalGaussianDistribution(object):
|
| 7 |
+
def __init__(self, parameters, upcast_fp32=True):
|
| 8 |
+
if upcast_fp32:
|
| 9 |
+
parameters = parameters.to(dtype=torch.float32)
|
| 10 |
+
|
| 11 |
+
self.parameters = parameters
|
| 12 |
+
self.mean, self.logvar = torch.chunk(parameters, 2, dim=1)
|
| 13 |
+
self.logvar = torch.clamp(self.logvar, -30.0, 20.0)
|
| 14 |
+
self.std = torch.exp(0.5 * self.logvar)
|
| 15 |
+
self.var = torch.exp(self.logvar)
|
| 16 |
+
|
| 17 |
+
@torch.compiler.disable
|
| 18 |
+
def sample(self, generator=None):
|
| 19 |
+
noise = torch.randn(self.mean.shape, generator=generator)
|
| 20 |
+
x = self.mean + self.std * noise.to(device=self.parameters.device)
|
| 21 |
+
return x
|
| 22 |
+
|
| 23 |
+
|
| 24 |
+
class ClsTokenAggregator:
|
| 25 |
+
def __init__(self, vae_model):
|
| 26 |
+
self.vae = vae_model
|
| 27 |
+
self.cls_tokens = []
|
| 28 |
+
|
| 29 |
+
def __enter__(self):
|
| 30 |
+
return self
|
| 31 |
+
|
| 32 |
+
def __exit__(self, exc_type, exc_val, exc_tb):
|
| 33 |
+
if self.cls_tokens and hasattr(self.vae.encoder, "loss_info"):
|
| 34 |
+
self.vae.encoder.loss_info["cls_token"] = torch.stack(
|
| 35 |
+
self.cls_tokens, dim=0
|
| 36 |
+
).mean(dim=0)
|
| 37 |
+
return False
|
| 38 |
+
|
| 39 |
+
def collect(self):
|
| 40 |
+
if (
|
| 41 |
+
hasattr(self.vae.encoder, "loss_info")
|
| 42 |
+
and "cls_token" in self.vae.encoder.loss_info
|
| 43 |
+
):
|
| 44 |
+
self.cls_tokens.append(self.vae.encoder.loss_info["cls_token"].clone())
|
| 45 |
+
|
| 46 |
+
def collect_stacked(self, num_tiles, sample_batch_size):
|
| 47 |
+
if (
|
| 48 |
+
hasattr(self.vae.encoder, "loss_info")
|
| 49 |
+
and "cls_token" in self.vae.encoder.loss_info
|
| 50 |
+
):
|
| 51 |
+
cls_token = self.vae.encoder.loss_info["cls_token"]
|
| 52 |
+
cls_token = cls_token.unflatten(0, (num_tiles, sample_batch_size))
|
| 53 |
+
self.cls_tokens.extend(token.clone() for token in cls_token)
|
FL2VA/video_vae/vae_processor.py
ADDED
|
@@ -0,0 +1,234 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 2 |
+
# Tensor pre/post-processing for the MiniMax H3 visual VAE.
|
| 3 |
+
import math
|
| 4 |
+
import numpy as np
|
| 5 |
+
import torch
|
| 6 |
+
from diffusers.utils import logging
|
| 7 |
+
from einops import rearrange
|
| 8 |
+
|
| 9 |
+
from .normalize import get_normalize_transform, get_denormalize_transform
|
| 10 |
+
|
| 11 |
+
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
| 12 |
+
|
| 13 |
+
|
| 14 |
+
class VAEProcessor:
|
| 15 |
+
|
| 16 |
+
def __init__(
|
| 17 |
+
self,
|
| 18 |
+
*,
|
| 19 |
+
vae_ratio,
|
| 20 |
+
vae_ratio_t,
|
| 21 |
+
clip_length,
|
| 22 |
+
frame_overlap,
|
| 23 |
+
token_overlap,
|
| 24 |
+
tokens_chunk_size,
|
| 25 |
+
isolated_last_frame,
|
| 26 |
+
latent_patch_size,
|
| 27 |
+
crop_mode,
|
| 28 |
+
pixel_norm_type="imagenet",
|
| 29 |
+
transform=None,
|
| 30 |
+
transform_rev=None,
|
| 31 |
+
use_3d_conv=False,
|
| 32 |
+
):
|
| 33 |
+
self.vae_ratio = vae_ratio
|
| 34 |
+
self.vae_ratio_t = vae_ratio_t
|
| 35 |
+
self.clip_length = clip_length
|
| 36 |
+
self.frame_overlap = frame_overlap
|
| 37 |
+
self.token_overlap = token_overlap
|
| 38 |
+
self.tokens_chunk_size = tokens_chunk_size
|
| 39 |
+
self.isolated_last_frame = isolated_last_frame
|
| 40 |
+
self.latent_patch_size = latent_patch_size
|
| 41 |
+
self.crop_mode = crop_mode
|
| 42 |
+
self.transform = transform or get_normalize_transform(pixel_norm_type)
|
| 43 |
+
self.transform_rev = transform_rev or get_denormalize_transform(pixel_norm_type)
|
| 44 |
+
self.use_3d_conv = use_3d_conv
|
| 45 |
+
|
| 46 |
+
def _ensure_list(self, data):
|
| 47 |
+
return data if isinstance(data, list) else [data]
|
| 48 |
+
|
| 49 |
+
def _align_to_total_patch_size(self, h, w):
|
| 50 |
+
total_patch_size = self.latent_patch_size * self.vae_ratio
|
| 51 |
+
new_h = (h // total_patch_size) * total_patch_size
|
| 52 |
+
new_w = (w // total_patch_size) * total_patch_size
|
| 53 |
+
return new_h, new_w
|
| 54 |
+
|
| 55 |
+
def _crop_to_align(self, tensor, new_h, new_w, is_video=False):
|
| 56 |
+
if is_video:
|
| 57 |
+
_, _, _, h, w = tensor.shape
|
| 58 |
+
else:
|
| 59 |
+
_, _, h, w = tensor.shape
|
| 60 |
+
|
| 61 |
+
if self.crop_mode == "center":
|
| 62 |
+
top = (h - new_h) // 2
|
| 63 |
+
left = (w - new_w) // 2
|
| 64 |
+
else:
|
| 65 |
+
top = 0
|
| 66 |
+
left = 0
|
| 67 |
+
|
| 68 |
+
if is_video:
|
| 69 |
+
return tensor[:, :, :, top : top + new_h, left : left + new_w]
|
| 70 |
+
else:
|
| 71 |
+
return tensor[:, :, top : top + new_h, left : left + new_w]
|
| 72 |
+
|
| 73 |
+
def _align_target_token(self, T, mode):
|
| 74 |
+
intra_tail = self.clip_length % self.vae_ratio_t
|
| 75 |
+
min_frames = intra_tail or self.vae_ratio_t
|
| 76 |
+
full_chunks = T // self.clip_length
|
| 77 |
+
remainder = T % self.clip_length
|
| 78 |
+
|
| 79 |
+
if remainder == 0:
|
| 80 |
+
return max(T, min_frames)
|
| 81 |
+
|
| 82 |
+
if mode == "pad":
|
| 83 |
+
aligned_r = (
|
| 84 |
+
math.ceil((remainder - intra_tail) / self.vae_ratio_t) * self.vae_ratio_t
|
| 85 |
+
+ intra_tail
|
| 86 |
+
)
|
| 87 |
+
if aligned_r > self.clip_length:
|
| 88 |
+
return (full_chunks + 1) * self.clip_length + intra_tail
|
| 89 |
+
return full_chunks * self.clip_length + aligned_r
|
| 90 |
+
else: # trim
|
| 91 |
+
k = (remainder - intra_tail) // self.vae_ratio_t
|
| 92 |
+
if k >= 0:
|
| 93 |
+
target = full_chunks * self.clip_length + k * self.vae_ratio_t + intra_tail
|
| 94 |
+
return max(target, min_frames)
|
| 95 |
+
elif full_chunks > 0:
|
| 96 |
+
return full_chunks * self.clip_length
|
| 97 |
+
else:
|
| 98 |
+
return min_frames
|
| 99 |
+
|
| 100 |
+
def _align_target(self, T, mode, granularity):
|
| 101 |
+
if granularity == "chunk":
|
| 102 |
+
step = self.clip_length
|
| 103 |
+
tail = self.frame_overlap
|
| 104 |
+
if self.isolated_last_frame:
|
| 105 |
+
tail += 1
|
| 106 |
+
|
| 107 |
+
k = math.ceil((T - tail) / step) if mode == "pad" else (T - tail) // step
|
| 108 |
+
return max(k, 1) * step + tail
|
| 109 |
+
|
| 110 |
+
isolated_extra = 1 if self.isolated_last_frame else 0
|
| 111 |
+
return self._align_target_token(T - isolated_extra, mode) + isolated_extra
|
| 112 |
+
|
| 113 |
+
def align_video_length(self, video_length, mode="pad", granularity="chunk"):
|
| 114 |
+
target = self._align_target(video_length, mode, granularity)
|
| 115 |
+
delta = target - video_length
|
| 116 |
+
if delta > 0 and mode == "trim":
|
| 117 |
+
raise ValueError(
|
| 118 |
+
f"Cannot trim {video_length} frames to valid length {target}: "
|
| 119 |
+
f"not enough frames (granularity={granularity})"
|
| 120 |
+
)
|
| 121 |
+
return delta
|
| 122 |
+
|
| 123 |
+
def align_video_length_2pass(self, video_length):
|
| 124 |
+
"""Return the leading/trailing frame pads and trailing latent drop.
|
| 125 |
+
|
| 126 |
+
This is the continuation-prefix (2-pass) alignment. The caller temporarily disables the model's normal token
|
| 127 |
+
drop and keeps these mirrored processor fields at zero.
|
| 128 |
+
"""
|
| 129 |
+
if self.isolated_last_frame:
|
| 130 |
+
raise ValueError(
|
| 131 |
+
"align_video_length_2pass does not support isolated_last_frame"
|
| 132 |
+
)
|
| 133 |
+
if self.token_overlap != 0 or self.frame_overlap != 0:
|
| 134 |
+
raise ValueError(
|
| 135 |
+
"align_video_length_2pass requires token_drop=0 alignment"
|
| 136 |
+
)
|
| 137 |
+
|
| 138 |
+
leading = self.align_video_length(
|
| 139 |
+
video_length, mode="pad", granularity="token"
|
| 140 |
+
)
|
| 141 |
+
token_aligned = video_length + leading
|
| 142 |
+
trailing = self.align_video_length(
|
| 143 |
+
token_aligned, mode="pad", granularity="chunk"
|
| 144 |
+
)
|
| 145 |
+
|
| 146 |
+
if trailing > 0:
|
| 147 |
+
intra_tail = self.clip_length % self.vae_ratio_t
|
| 148 |
+
full_chunks = token_aligned // self.clip_length
|
| 149 |
+
remainder = token_aligned % self.clip_length
|
| 150 |
+
real_tokens = full_chunks * self.tokens_chunk_size
|
| 151 |
+
if remainder > 0:
|
| 152 |
+
real_tokens += (
|
| 153 |
+
(remainder - intra_tail) // self.vae_ratio_t + 1
|
| 154 |
+
)
|
| 155 |
+
drop_tokens = (
|
| 156 |
+
self.get_latent_length(token_aligned + trailing) - real_tokens
|
| 157 |
+
)
|
| 158 |
+
else:
|
| 159 |
+
drop_tokens = 0
|
| 160 |
+
|
| 161 |
+
return leading, trailing, drop_tokens
|
| 162 |
+
|
| 163 |
+
def get_suitable_video_length(self, video_length, verbose=False):
|
| 164 |
+
used_frame_length = video_length + self.align_video_length(
|
| 165 |
+
video_length, mode="trim", granularity="chunk"
|
| 166 |
+
)
|
| 167 |
+
if verbose:
|
| 168 |
+
logger.info(
|
| 169 |
+
f"Pick first {used_frame_length} frames from {video_length}-frame video"
|
| 170 |
+
)
|
| 171 |
+
return used_frame_length
|
| 172 |
+
|
| 173 |
+
def get_latent_length(self, video_length):
|
| 174 |
+
tail_frame = self.frame_overlap
|
| 175 |
+
tail_token = self.token_overlap
|
| 176 |
+
if self.isolated_last_frame:
|
| 177 |
+
tail_frame += 1
|
| 178 |
+
tail_token += 1
|
| 179 |
+
|
| 180 |
+
video_length = self.get_suitable_video_length(video_length)
|
| 181 |
+
latent_length = (
|
| 182 |
+
int((video_length - tail_frame) // self.clip_length)
|
| 183 |
+
* self.tokens_chunk_size
|
| 184 |
+
+ tail_token
|
| 185 |
+
)
|
| 186 |
+
return latent_length
|
| 187 |
+
|
| 188 |
+
|
| 189 |
+
|
| 190 |
+
def transform_tensor(self, tensor):
|
| 191 |
+
B, T = None, None
|
| 192 |
+
if tensor.ndim == 5:
|
| 193 |
+
if tensor.shape[2] == 3:
|
| 194 |
+
tensor = tensor.transpose(1, 2)
|
| 195 |
+
B, _, T, _, _ = tensor.shape
|
| 196 |
+
tensor = rearrange(tensor, "b c t h w -> (b t) c h w")
|
| 197 |
+
elif tensor.ndim == 4:
|
| 198 |
+
if tensor.shape[0] == 3:
|
| 199 |
+
tensor = tensor.transpose(0, 1)
|
| 200 |
+
elif tensor.ndim == 3:
|
| 201 |
+
tensor = tensor.unsqueeze(0)
|
| 202 |
+
else:
|
| 203 |
+
raise ValueError(f"Unsupported tensor shape: {tensor.shape}")
|
| 204 |
+
|
| 205 |
+
tensor = self.transform(tensor)
|
| 206 |
+
|
| 207 |
+
if B is not None and T is not None:
|
| 208 |
+
tensor = rearrange(tensor, "(b t) c h w -> b c t h w", b=B, t=T)
|
| 209 |
+
|
| 210 |
+
return tensor.contiguous()
|
| 211 |
+
|
| 212 |
+
def revert_tensor(self, tensor):
|
| 213 |
+
B, T = None, None
|
| 214 |
+
if self.use_3d_conv:
|
| 215 |
+
tensor = tensor.unsqueeze(2) if tensor.ndim == 4 else tensor
|
| 216 |
+
B, _, T, _, _ = tensor.shape
|
| 217 |
+
tensor = rearrange(tensor, "b c t h w -> (b t) c h w")
|
| 218 |
+
tensor_rev = self.transform_rev(tensor).clamp(0, 1)
|
| 219 |
+
if B is not None:
|
| 220 |
+
tensor_rev = rearrange(tensor_rev, "(b t) c h w -> b c t h w", b=B, t=T)
|
| 221 |
+
return tensor_rev.contiguous()
|
| 222 |
+
|
| 223 |
+
@staticmethod
|
| 224 |
+
def convert_numpy_to_tensor(numpy_array, device=None):
|
| 225 |
+
if isinstance(numpy_array, list):
|
| 226 |
+
numpy_array = np.stack(numpy_array, axis=0)
|
| 227 |
+
numpy_array = numpy_array.astype(np.float32)
|
| 228 |
+
tensor = torch.from_numpy(numpy_array)
|
| 229 |
+
tensor = tensor.permute(0, 3, 1, 2)
|
| 230 |
+
tensor = tensor / 255.0
|
| 231 |
+
if device is not None:
|
| 232 |
+
tensor = tensor.to(device)
|
| 233 |
+
return tensor
|
| 234 |
+
|
FL2VA/video_vae/vae_vit.py
ADDED
|
@@ -0,0 +1,380 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# SPDX-License-Identifier: Apache-2.0
|
| 2 |
+
# ViT3D decoder for the MiniMax H3 visual VAE (inference-only bundle).
|
| 3 |
+
import torch
|
| 4 |
+
import torch.nn as nn
|
| 5 |
+
import torch.distributed as dist
|
| 6 |
+
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
| 7 |
+
from diffusers.models.modeling_utils import ModelMixin
|
| 8 |
+
from diffusers.utils import logging
|
| 9 |
+
|
| 10 |
+
from .attention import maybe_checkpoint
|
| 11 |
+
from .base_module import TransformerBlock, RotaryEmbeddingND
|
| 12 |
+
from .flash import make_block_causal_mask_mod
|
| 13 |
+
from .func import create_token_ids
|
| 14 |
+
from .parallel import get_subseq, gather_subseq, get_parallel_state
|
| 15 |
+
|
| 16 |
+
logger = logging.get_logger(__name__)
|
| 17 |
+
|
| 18 |
+
|
| 19 |
+
def _linear_with_module_dtype(linear, tensor, out_dtype=None):
|
| 20 |
+
weight = getattr(linear, "weight", None)
|
| 21 |
+
target_dtype = getattr(weight, "dtype", tensor.dtype)
|
| 22 |
+
output = linear(tensor.to(target_dtype))
|
| 23 |
+
if out_dtype is not None and output.dtype != out_dtype:
|
| 24 |
+
output = output.to(out_dtype)
|
| 25 |
+
return output
|
| 26 |
+
|
| 27 |
+
|
| 28 |
+
def _make_seq_len_mask_mod(seq_len, base_mask_mod=None):
|
| 29 |
+
if base_mask_mod is None:
|
| 30 |
+
def mask_mod(b, h, q_idx, kv_idx, seqlen_info, aux_tensors):
|
| 31 |
+
return (q_idx < seq_len) & (kv_idx < seq_len)
|
| 32 |
+
|
| 33 |
+
mask_mod.block_sparse_cache_key = ("seq_len", seq_len)
|
| 34 |
+
return mask_mod
|
| 35 |
+
|
| 36 |
+
def mask_mod(b, h, q_idx, kv_idx, seqlen_info, aux_tensors):
|
| 37 |
+
return (
|
| 38 |
+
(q_idx < seq_len)
|
| 39 |
+
& (kv_idx < seq_len)
|
| 40 |
+
& base_mask_mod(b, h, q_idx, kv_idx, seqlen_info, aux_tensors)
|
| 41 |
+
)
|
| 42 |
+
|
| 43 |
+
base_cache_key = getattr(base_mask_mod, "block_sparse_cache_key", None)
|
| 44 |
+
if base_cache_key is not None:
|
| 45 |
+
mask_mod.block_sparse_cache_key = ("seq_len", seq_len, base_cache_key)
|
| 46 |
+
if hasattr(base_mask_mod, "use_fast_sampling"):
|
| 47 |
+
mask_mod.use_fast_sampling = base_mask_mod.use_fast_sampling
|
| 48 |
+
return mask_mod
|
| 49 |
+
|
| 50 |
+
|
| 51 |
+
|
| 52 |
+
|
| 53 |
+
|
| 54 |
+
|
| 55 |
+
def _pack_tensors_3d(tensors, patch_size, patch_size_t):
|
| 56 |
+
batch_size, num_channels_tensors, temporal, height, width = tensors.shape
|
| 57 |
+
|
| 58 |
+
tensors = tensors.view(
|
| 59 |
+
batch_size,
|
| 60 |
+
num_channels_tensors,
|
| 61 |
+
temporal // patch_size_t,
|
| 62 |
+
patch_size_t,
|
| 63 |
+
height // patch_size,
|
| 64 |
+
patch_size,
|
| 65 |
+
width // patch_size,
|
| 66 |
+
patch_size,
|
| 67 |
+
)
|
| 68 |
+
tensors = tensors.permute(0, 2, 4, 6, 1, 3, 5, 7)
|
| 69 |
+
tensors = tensors.reshape(
|
| 70 |
+
batch_size,
|
| 71 |
+
(temporal // patch_size_t) * (height // patch_size) * (width // patch_size),
|
| 72 |
+
num_channels_tensors * patch_size_t * patch_size * patch_size,
|
| 73 |
+
)
|
| 74 |
+
return tensors
|
| 75 |
+
|
| 76 |
+
|
| 77 |
+
def _unpack_tensors_3d(tensors, patch_size, patch_size_t, temporal, height, width):
|
| 78 |
+
batch_size, num_patches, channels = tensors.shape
|
| 79 |
+
num_channels_tensors = channels // (patch_size_t * patch_size * patch_size)
|
| 80 |
+
|
| 81 |
+
tensors = tensors.view(
|
| 82 |
+
batch_size,
|
| 83 |
+
temporal // patch_size_t,
|
| 84 |
+
height // patch_size,
|
| 85 |
+
width // patch_size,
|
| 86 |
+
num_channels_tensors,
|
| 87 |
+
patch_size_t,
|
| 88 |
+
patch_size,
|
| 89 |
+
patch_size,
|
| 90 |
+
)
|
| 91 |
+
tensors = tensors.permute(0, 4, 1, 5, 2, 6, 3, 7).contiguous()
|
| 92 |
+
tensors = tensors.reshape(batch_size, num_channels_tensors, temporal, height, width)
|
| 93 |
+
return tensors
|
| 94 |
+
|
| 95 |
+
|
| 96 |
+
class ViTBase(ModelMixin, ConfigMixin):
|
| 97 |
+
"""Base class for ViT Encoder and Decoder with common functionality."""
|
| 98 |
+
|
| 99 |
+
_supports_gradient_checkpointing = True
|
| 100 |
+
_no_split_modules = ["TransformerBlock"]
|
| 101 |
+
gradient_checkpointing_mode = "full"
|
| 102 |
+
|
| 103 |
+
def _set_gradient_checkpointing(self, module, value=False):
|
| 104 |
+
if hasattr(module, "gradient_checkpointing"):
|
| 105 |
+
module.gradient_checkpointing = value
|
| 106 |
+
|
| 107 |
+
def set_spatial_parallel(self, enabled):
|
| 108 |
+
self.spatial_parallel = enabled
|
| 109 |
+
if hasattr(self, "transformer_blocks"):
|
| 110 |
+
for block in self.transformer_blocks:
|
| 111 |
+
block.attn.spatial_parallel = enabled
|
| 112 |
+
|
| 113 |
+
def _init_weights(self):
|
| 114 |
+
def basic_init(m):
|
| 115 |
+
if isinstance(m, nn.Linear):
|
| 116 |
+
nn.init.xavier_uniform_(m.weight)
|
| 117 |
+
if m.bias is not None:
|
| 118 |
+
nn.init.constant_(m.bias, 0)
|
| 119 |
+
|
| 120 |
+
self.apply(basic_init)
|
| 121 |
+
|
| 122 |
+
def init_mask_config(self, dim, is_3d=False):
|
| 123 |
+
self._mask_dim = dim
|
| 124 |
+
self._mask_is_3d = is_3d
|
| 125 |
+
self.register_buffer("mask_token", torch.zeros(1, 1, dim))
|
| 126 |
+
|
| 127 |
+
def set_mask_config(self, mask_config):
|
| 128 |
+
self.mask_prob = mask_config.get("mask_prob", 0.0)
|
| 129 |
+
self.mask_enabled = self.mask_prob > 0
|
| 130 |
+
self.mask_style = mask_config.get("mask_style", "replace")
|
| 131 |
+
if self.mask_enabled and self.mask_style == "drop" and self.mask_prob < 1.0:
|
| 132 |
+
logger.warning("mask_style='drop' with mask_prob < 1.0")
|
| 133 |
+
if self._mask_is_3d:
|
| 134 |
+
self.temporal_scale_range = mask_config.get("temporal_scale_range", (0.3, 0.5))
|
| 135 |
+
self.spatial_scale_range = mask_config.get("spatial_scale_range", (0.1, 0.25))
|
| 136 |
+
self.min_mask_ratio = mask_config.get("min_mask_ratio", 0.75)
|
| 137 |
+
self.max_mask_ratio = mask_config.get("max_mask_ratio", 0.95)
|
| 138 |
+
else:
|
| 139 |
+
self.spatial_scale_range = mask_config.get("spatial_scale_range", (0.15, 0.15))
|
| 140 |
+
self.min_mask_ratio = mask_config.get("min_mask_ratio", 0.5)
|
| 141 |
+
self.max_mask_ratio = mask_config.get("max_mask_ratio", 0.75)
|
| 142 |
+
self.aspect_ratio_range = mask_config.get("aspect_ratio_range", (0.75, 1.5))
|
| 143 |
+
self.max_retries = mask_config.get("max_retries", 100)
|
| 144 |
+
if self.mask_enabled and self.mask_style == "drop" and getattr(self, "t_causal", False):
|
| 145 |
+
logger.warning("mask_style='drop' with t_causal may cause issues")
|
| 146 |
+
if self.mask_enabled and "mask_token" in self._buffers:
|
| 147 |
+
del self._buffers["mask_token"]
|
| 148 |
+
self.mask_token = nn.Parameter(torch.randn(1, 1, self._mask_dim) * 0.02)
|
| 149 |
+
|
| 150 |
+
def init_suffix_tokens(self, dim, num_register_tokens, has_cls_token=True):
|
| 151 |
+
self.num_register_tokens = num_register_tokens
|
| 152 |
+
if num_register_tokens > 0:
|
| 153 |
+
self.register_tokens = nn.Parameter(torch.randn(1, num_register_tokens, dim) * 0.02)
|
| 154 |
+
else:
|
| 155 |
+
self.register_tokens = None
|
| 156 |
+
if has_cls_token:
|
| 157 |
+
self.cls_token = nn.Parameter(torch.randn(1, 1, dim) * 0.02)
|
| 158 |
+
|
| 159 |
+
def apply_mask_preprocess(self, hidden_states, img_ids, patch_dims, num_suffix):
|
| 160 |
+
if self.training and self.mask_enabled:
|
| 161 |
+
raise NotImplementedError(
|
| 162 |
+
"mask modeling is not supported in this inference-only bundle"
|
| 163 |
+
)
|
| 164 |
+
return hidden_states, img_ids
|
| 165 |
+
|
| 166 |
+
def forward_transformer_blocks(self, hidden_states, rotary_pos_emb, pack_info=None):
|
| 167 |
+
if pack_info is None:
|
| 168 |
+
pack_info = {}
|
| 169 |
+
for block in self.transformer_blocks:
|
| 170 |
+
hidden_states = maybe_checkpoint(
|
| 171 |
+
self, block, hidden_states, rotary_pos_emb, pack_info
|
| 172 |
+
)
|
| 173 |
+
return hidden_states
|
| 174 |
+
|
| 175 |
+
def _pad_for_sp(self, hidden_states, img_ids, pack_info=None):
|
| 176 |
+
if pack_info is None:
|
| 177 |
+
pack_info = {}
|
| 178 |
+
if not self.spatial_parallel:
|
| 179 |
+
return hidden_states, img_ids, pack_info, 0
|
| 180 |
+
|
| 181 |
+
seq_len = hidden_states.shape[1]
|
| 182 |
+
sp_size = get_parallel_state().get("sp_size", 1)
|
| 183 |
+
pad_len = (-seq_len) % sp_size
|
| 184 |
+
if pad_len == 0:
|
| 185 |
+
return hidden_states, img_ids, pack_info, 0
|
| 186 |
+
|
| 187 |
+
hidden_states = torch.nn.functional.pad(hidden_states, (0, 0, 0, pad_len))
|
| 188 |
+
img_ids = torch.nn.functional.pad(img_ids, (0, 0, 0, pad_len))
|
| 189 |
+
|
| 190 |
+
pack_info = dict(pack_info)
|
| 191 |
+
base_mask_mod = pack_info.get("mask_mod")
|
| 192 |
+
pack_info["mask_mod"] = _make_seq_len_mask_mod(seq_len, base_mask_mod)
|
| 193 |
+
pack_info.pop("block_sparse", None)
|
| 194 |
+
return hidden_states, img_ids, pack_info, pad_len
|
| 195 |
+
|
| 196 |
+
@staticmethod
|
| 197 |
+
def _unpad_for_sp(hidden_states, pad_len):
|
| 198 |
+
if pad_len == 0:
|
| 199 |
+
return hidden_states
|
| 200 |
+
return hidden_states[:, :-pad_len, :]
|
| 201 |
+
|
| 202 |
+
def apply_mask_postprocess(self, hidden_states, num_patches):
|
| 203 |
+
if self.training and self.mask_enabled and self.mask_style == "drop":
|
| 204 |
+
raise NotImplementedError(
|
| 205 |
+
"mask modeling is not supported in this inference-only bundle"
|
| 206 |
+
)
|
| 207 |
+
return hidden_states
|
| 208 |
+
|
| 209 |
+
|
| 210 |
+
|
| 211 |
+
|
| 212 |
+
|
| 213 |
+
|
| 214 |
+
|
| 215 |
+
|
| 216 |
+
class ViT3DDecoder(ViTBase):
|
| 217 |
+
"""Vision Transformer Video Decoder using TransformerBlock."""
|
| 218 |
+
|
| 219 |
+
@register_to_config
|
| 220 |
+
def __init__(
|
| 221 |
+
self,
|
| 222 |
+
patch_size: int = 16,
|
| 223 |
+
patch_size_t: int = 4,
|
| 224 |
+
t_causal: bool = False,
|
| 225 |
+
in_channels: int = 16,
|
| 226 |
+
out_channels: int = 3,
|
| 227 |
+
num_layers: int = 24,
|
| 228 |
+
heads: int = 16,
|
| 229 |
+
dim_head: int = 64,
|
| 230 |
+
norm_type: str = "layer_norm",
|
| 231 |
+
norm_affine: bool = True,
|
| 232 |
+
qk_norm_type: str = None,
|
| 233 |
+
qk_norm_affine: bool = False,
|
| 234 |
+
ffn_activation_fn: str = "gelu",
|
| 235 |
+
ffn_use_gated: bool = False,
|
| 236 |
+
rope_theta: float = 100.0,
|
| 237 |
+
rope_dim_ratio: float = 1.0,
|
| 238 |
+
bias: bool = True,
|
| 239 |
+
eps: float = 1e-5,
|
| 240 |
+
num_register_tokens: int = 4,
|
| 241 |
+
mask_config: dict = {},
|
| 242 |
+
**kwargs,
|
| 243 |
+
):
|
| 244 |
+
super().__init__()
|
| 245 |
+
|
| 246 |
+
dim = heads * dim_head
|
| 247 |
+
rope_apply_dim = int(dim_head * rope_dim_ratio)
|
| 248 |
+
|
| 249 |
+
self.pos_embed = RotaryEmbeddingND(rope_apply_dim, rope_theta, n_dim=3, use_angle=True)
|
| 250 |
+
|
| 251 |
+
self.x_embedder = nn.Linear(in_channels, dim)
|
| 252 |
+
|
| 253 |
+
self.init_suffix_tokens(dim, num_register_tokens, has_cls_token=False)
|
| 254 |
+
|
| 255 |
+
self.t_causal = t_causal
|
| 256 |
+
|
| 257 |
+
self.transformer_blocks = nn.ModuleList(
|
| 258 |
+
[
|
| 259 |
+
TransformerBlock(
|
| 260 |
+
heads=heads,
|
| 261 |
+
dim_head=dim_head,
|
| 262 |
+
norm_type=norm_type,
|
| 263 |
+
norm_affine=norm_affine,
|
| 264 |
+
qk_norm_type=qk_norm_type,
|
| 265 |
+
qk_norm_affine=qk_norm_affine,
|
| 266 |
+
ffn_activation_fn=ffn_activation_fn,
|
| 267 |
+
ffn_use_gated=ffn_use_gated,
|
| 268 |
+
bias=bias,
|
| 269 |
+
eps=eps,
|
| 270 |
+
**kwargs,
|
| 271 |
+
)
|
| 272 |
+
for _ in range(num_layers)
|
| 273 |
+
]
|
| 274 |
+
)
|
| 275 |
+
|
| 276 |
+
self.spatial_parallel = False
|
| 277 |
+
for block in self.transformer_blocks:
|
| 278 |
+
block.attn.spatial_parallel = False
|
| 279 |
+
|
| 280 |
+
self.norm_out = nn.LayerNorm(dim, elementwise_affine=norm_affine, eps=eps)
|
| 281 |
+
patch_dim = out_channels * patch_size_t * patch_size * patch_size
|
| 282 |
+
self.proj_out = nn.Linear(dim, patch_dim)
|
| 283 |
+
|
| 284 |
+
self.init_mask_config(dim, is_3d=True)
|
| 285 |
+
self.set_mask_config(mask_config)
|
| 286 |
+
|
| 287 |
+
self._init_weights()
|
| 288 |
+
self.gradient_checkpointing = False
|
| 289 |
+
|
| 290 |
+
if len(kwargs) > 0 and (not dist.is_initialized() or dist.get_rank() == 0):
|
| 291 |
+
logger.warning(f"Unused kwargs: {kwargs}")
|
| 292 |
+
|
| 293 |
+
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 294 |
+
self.loss_info = {}
|
| 295 |
+
|
| 296 |
+
B, C, latent_T, latent_H, latent_W = x.shape
|
| 297 |
+
patch_size = self.config.patch_size
|
| 298 |
+
patch_size_t = self.config.patch_size_t
|
| 299 |
+
num_suffix = 1 + self.num_register_tokens
|
| 300 |
+
|
| 301 |
+
hidden_states = _pack_tensors_3d(x, 1, 1)
|
| 302 |
+
latent_size = (latent_T, latent_H, latent_W)
|
| 303 |
+
|
| 304 |
+
with torch.autocast("cuda", enabled=False):
|
| 305 |
+
hidden_states = _linear_with_module_dtype(self.x_embedder, hidden_states, hidden_states.dtype)
|
| 306 |
+
|
| 307 |
+
num_patches = hidden_states.shape[1]
|
| 308 |
+
|
| 309 |
+
tokens = [hidden_states]
|
| 310 |
+
|
| 311 |
+
if self.register_tokens is not None:
|
| 312 |
+
register_tokens = self.register_tokens.expand(B, -1, -1)
|
| 313 |
+
tokens.append(register_tokens)
|
| 314 |
+
|
| 315 |
+
cls_token = torch.zeros_like(hidden_states[:, 0:1, :])
|
| 316 |
+
tokens.append(cls_token)
|
| 317 |
+
hidden_states = torch.cat(tokens, dim=1)
|
| 318 |
+
|
| 319 |
+
patch_dims = [latent_T, latent_H, latent_W]
|
| 320 |
+
img_ids = create_token_ids(latent_size, x.device, x.dtype).expand(B, -1, -1)
|
| 321 |
+
suffix_ids = torch.zeros((B, num_suffix, 3), device=x.device, dtype=img_ids.dtype)
|
| 322 |
+
img_ids = torch.cat([img_ids, suffix_ids], dim=1)
|
| 323 |
+
|
| 324 |
+
hidden_states, img_ids = self.apply_mask_preprocess(hidden_states, img_ids, patch_dims, num_suffix)
|
| 325 |
+
|
| 326 |
+
pack_info = {}
|
| 327 |
+
if self.t_causal:
|
| 328 |
+
spatial_size = latent_H * latent_W
|
| 329 |
+
mask_mod = make_block_causal_mask_mod(
|
| 330 |
+
num_tokens=num_patches,
|
| 331 |
+
block_size=spatial_size,
|
| 332 |
+
suffix=True,
|
| 333 |
+
)
|
| 334 |
+
pack_info["mask_mod"] = mask_mod
|
| 335 |
+
|
| 336 |
+
hidden_states, img_ids, pack_info, sp_pad_len = self._pad_for_sp(hidden_states, img_ids, pack_info)
|
| 337 |
+
|
| 338 |
+
rotary_pos_emb = self.pos_embed(img_ids)
|
| 339 |
+
|
| 340 |
+
if self.spatial_parallel:
|
| 341 |
+
hidden_states = get_subseq(hidden_states)
|
| 342 |
+
|
| 343 |
+
for block in self.transformer_blocks:
|
| 344 |
+
hidden_states = maybe_checkpoint(
|
| 345 |
+
self, block, hidden_states, rotary_pos_emb, pack_info
|
| 346 |
+
)
|
| 347 |
+
|
| 348 |
+
if self.spatial_parallel:
|
| 349 |
+
hidden_states = gather_subseq(hidden_states)
|
| 350 |
+
hidden_states = self._unpad_for_sp(hidden_states, sp_pad_len)
|
| 351 |
+
|
| 352 |
+
hidden_states = self.norm_out(hidden_states)
|
| 353 |
+
|
| 354 |
+
hidden_states = self.apply_mask_postprocess(hidden_states, num_patches)
|
| 355 |
+
|
| 356 |
+
with torch.autocast("cuda", enabled=False):
|
| 357 |
+
output = _linear_with_module_dtype(self.proj_out, hidden_states, hidden_states.dtype)
|
| 358 |
+
|
| 359 |
+
output = output[:, :num_patches, :]
|
| 360 |
+
|
| 361 |
+
video_t = latent_size[0] * patch_size_t
|
| 362 |
+
video_h = latent_size[1] * patch_size
|
| 363 |
+
video_w = latent_size[2] * patch_size
|
| 364 |
+
output = _unpack_tensors_3d(output, patch_size, patch_size_t, video_t, video_h, video_w)
|
| 365 |
+
|
| 366 |
+
return output
|
| 367 |
+
|
| 368 |
+
|
| 369 |
+
|
| 370 |
+
|
| 371 |
+
|
| 372 |
+
|
| 373 |
+
|
| 374 |
+
|
| 375 |
+
|
| 376 |
+
|
| 377 |
+
|
| 378 |
+
|
| 379 |
+
|
| 380 |
+
|
Ref2VA/audio_vae/model.safetensors
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:37dddc2f3e6d5d5139d823d5ea283bbf304dadcb885b1ccda818aa13dade5ea2
|
| 3 |
+
size 605429308
|
assets/fl2va.mp4
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:5a5d6e0fcdda7e6a5d98cb92d9e5cb4b6b0b44d6e1dfd5a370d16dfaa79b9ae6
|
| 3 |
+
size 1064614
|
assets/full-arch.png
ADDED
|
Git LFS Details
|
assets/h3_direct_2k.mp4
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:cc928dc5f406ca22d66656f98717808fd16ede2bfb5a3483c9b71264a9d74678
|
| 3 |
+
size 8946498
|
assets/h3_direct_768p.mp4
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:642e444e8cfa9d3417da69f429598d067bb8123639bbbe462de5687b583ac641
|
| 3 |
+
size 2385173
|
assets/i2va.mp4
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:5a5d6e0fcdda7e6a5d98cb92d9e5cb4b6b0b44d6e1dfd5a370d16dfaa79b9ae6
|
| 3 |
+
size 1064614
|
assets/i2va_2k.mp4
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:dcfec002e432208a9feeeb90b7867a00c1d6a18ce4707fa94df0930997802840
|
| 3 |
+
size 3334690
|
assets/i2va_direct_2k.mp4
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:c3c1eef07982164c7be26a2c8286c14a3f3d15afca94006856d3501873cc7fed
|
| 3 |
+
size 4569735
|
assets/i2va_direct_768p.mp4
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:fc73d7633bed7f0125b9219ae97e9f7c75c26176aae82bbb4c5d1576187d96f1
|
| 3 |
+
size 1448158
|
assets/minimax-h3.png
ADDED
|
Git LFS Details
|
assets/overview.png
ADDED
|
Git LFS Details
|
assets/r2va.mp4
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:5c9d33dfc517438c59e93da7b3d0b2a62e3cf69bc2adb7d1b2d33d7d8df7161f
|
| 3 |
+
size 850734
|
assets/r2va_2k.mp4
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:c3c96ce2bcd22b61ca65974c1a071b1659398644416cae2609f1ba026bd3ce35
|
| 3 |
+
size 2788329
|
assets/r2va_direct_2k.mp4
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:26f07ac79f8731360790ae1f046ab9d78ba0b3ba5fe84f61578696856bbec993
|
| 3 |
+
size 3134612
|