yuangai commited on
Commit
849926f
·
1 Parent(s): e45c5c8

init space

Browse files
.gitignore ADDED
@@ -0,0 +1,16 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ wandb
2
+ __pycache__
3
+ .vscode
4
+ notebooks
5
+ results
6
+ *.ipynb
7
+ *.ipynb_checkpoints
8
+ eval_results
9
+ tests
10
+ downloads
11
+ ckpts
12
+ demo_images
13
+ eval/OneIG-Benchmark/models
14
+ eval/OneIG-Benchmark/scripts/style/models
15
+ eval/OneIG-Benchmark/results*
16
+ models/
.gradio/certificate.pem ADDED
@@ -0,0 +1,31 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ -----BEGIN CERTIFICATE-----
2
+ MIIFazCCA1OgAwIBAgIRAIIQz7DSQONZRGPgu2OCiwAwDQYJKoZIhvcNAQELBQAw
3
+ TzELMAkGA1UEBhMCVVMxKTAnBgNVBAoTIEludGVybmV0IFNlY3VyaXR5IFJlc2Vh
4
+ cmNoIEdyb3VwMRUwEwYDVQQDEwxJU1JHIFJvb3QgWDEwHhcNMTUwNjA0MTEwNDM4
5
+ WhcNMzUwNjA0MTEwNDM4WjBPMQswCQYDVQQGEwJVUzEpMCcGA1UEChMgSW50ZXJu
6
+ ZXQgU2VjdXJpdHkgUmVzZWFyY2ggR3JvdXAxFTATBgNVBAMTDElTUkcgUm9vdCBY
7
+ MTCCAiIwDQYJKoZIhvcNAQEBBQADggIPADCCAgoCggIBAK3oJHP0FDfzm54rVygc
8
+ h77ct984kIxuPOZXoHj3dcKi/vVqbvYATyjb3miGbESTtrFj/RQSa78f0uoxmyF+
9
+ 0TM8ukj13Xnfs7j/EvEhmkvBioZxaUpmZmyPfjxwv60pIgbz5MDmgK7iS4+3mX6U
10
+ A5/TR5d8mUgjU+g4rk8Kb4Mu0UlXjIB0ttov0DiNewNwIRt18jA8+o+u3dpjq+sW
11
+ T8KOEUt+zwvo/7V3LvSye0rgTBIlDHCNAymg4VMk7BPZ7hm/ELNKjD+Jo2FR3qyH
12
+ B5T0Y3HsLuJvW5iB4YlcNHlsdu87kGJ55tukmi8mxdAQ4Q7e2RCOFvu396j3x+UC
13
+ B5iPNgiV5+I3lg02dZ77DnKxHZu8A/lJBdiB3QW0KtZB6awBdpUKD9jf1b0SHzUv
14
+ KBds0pjBqAlkd25HN7rOrFleaJ1/ctaJxQZBKT5ZPt0m9STJEadao0xAH0ahmbWn
15
+ OlFuhjuefXKnEgV4We0+UXgVCwOPjdAvBbI+e0ocS3MFEvzG6uBQE3xDk3SzynTn
16
+ jh8BCNAw1FtxNrQHusEwMFxIt4I7mKZ9YIqioymCzLq9gwQbooMDQaHWBfEbwrbw
17
+ qHyGO0aoSCqI3Haadr8faqU9GY/rOPNk3sgrDQoo//fb4hVC1CLQJ13hef4Y53CI
18
+ rU7m2Ys6xt0nUW7/vGT1M0NPAgMBAAGjQjBAMA4GA1UdDwEB/wQEAwIBBjAPBgNV
19
+ HRMBAf8EBTADAQH/MB0GA1UdDgQWBBR5tFnme7bl5AFzgAiIyBpY9umbbjANBgkq
20
+ hkiG9w0BAQsFAAOCAgEAVR9YqbyyqFDQDLHYGmkgJykIrGF1XIpu+ILlaS/V9lZL
21
+ ubhzEFnTIZd+50xx+7LSYK05qAvqFyFWhfFQDlnrzuBZ6brJFe+GnY+EgPbk6ZGQ
22
+ 3BebYhtF8GaV0nxvwuo77x/Py9auJ/GpsMiu/X1+mvoiBOv/2X/qkSsisRcOj/KK
23
+ NFtY2PwByVS5uCbMiogziUwthDyC3+6WVwW6LLv3xLfHTjuCvjHIInNzktHCgKQ5
24
+ ORAzI4JMPJ+GslWYHb4phowim57iaztXOoJwTdwJx4nLCgdNbOhdjsnvzqvHu7Ur
25
+ TkXWStAmzOVyyghqpZXjFaH3pO3JLF+l+/+sKAIuvtd7u+Nxe5AW0wdeRlN8NwdC
26
+ jNPElpzVmbUq4JUagEiuTDkHzsxHpFKVK7q4+63SM1N95R1NbdWhscdCb+ZAJzVc
27
+ oyi3B43njTOQ5yOf+1CceWxG1bQVs5ZufpsMljq4Ui0/1lvh+wjChP4kqKOJ2qxq
28
+ 4RgqsahDYVvTH9w7jXbyLeiNdd8XM2w9U/t7y0Ff/9yi0GE44Za4rF2LN9d11TPA
29
+ mRGunUHBcnWEvgJBQl9nJEiU0Zsnvgc/ubhPgXRR4Xq37Z0j4r7g1SgEEzwxA57d
30
+ emyPxgcYxn/eR44/KJ4EBs+lVDR3veyJm+kXQ99b21/+jh5Xos1AnX5iItreGCc=
31
+ -----END CERTIFICATE-----
app.py ADDED
@@ -0,0 +1,231 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import logging
2
+ import os
3
+ import random
4
+ import re
5
+ import sys
6
+ import warnings
7
+ import gradio as gr
8
+ import math
9
+ import torch
10
+ import subprocess
11
+ subprocess.run(
12
+ "pip install flash-attn==2.8.2 --no-build-isolation",
13
+ env={"FLASH_ATTENTION_SKIP_CUDA_BUILD": "TRUE"},
14
+ shell=True,
15
+ )
16
+
17
+ sys.path.append(os.path.dirname(os.path.abspath(__file__)))
18
+
19
+ try:
20
+ from modeling.t2i_pipeline import BitDanceT2IPipeline
21
+ except ImportError:
22
+ print("Warning: Could not import BitDanceT2IPipeline. Please ensure 'modeling' folder is present.")
23
+
24
+ from huggingface_hub import snapshot_download
25
+
26
+ save_dir = "models/BitDance-14B-64x"
27
+ repo_id = "shallowdream204/BitDance-14B-64x"
28
+ cache_dir = save_dir + "/cache"
29
+
30
+ snapshot_download(cache_dir=cache_dir,
31
+ local_dir=save_dir,
32
+ repo_id=repo_id,
33
+ local_dir_use_symlinks=False,
34
+ resume_download=True,
35
+ allow_patterns=["*.json", "*.safetensors", "*.bin", "*.py", "*.md", "*.txt"],
36
+ )
37
+
38
+ # ==================== Environment Variables ==================================
39
+ MODEL_PATH = save_dir
40
+
41
+ # =============================================================================
42
+ warnings.filterwarnings("ignore")
43
+ logging.getLogger("transformers").setLevel(logging.ERROR)
44
+
45
+ # ==================== Resolution Settings ====================================
46
+ RAW_RESOLUTIONS = [
47
+ [2048, 512],
48
+ [1920, 512],
49
+ [1536, 640],
50
+ [1280, 768],
51
+ [1152, 896],
52
+ [1024, 1024],
53
+ [896, 1152],
54
+ [768, 1280],
55
+ [640, 1536],
56
+ [512, 1920],
57
+ [512, 2048],
58
+ ]
59
+
60
+ RESOLUTION_CHOICES = []
61
+ for w, h in RAW_RESOLUTIONS:
62
+ divisor = math.gcd(w, h)
63
+ ratio_w = w // divisor
64
+ ratio_h = h // divisor
65
+ label = f"{w}x{h} ({ratio_w}:{ratio_h})"
66
+ RESOLUTION_CHOICES.append(label)
67
+
68
+ DEFAULT_RES = "1024x1024 (1:1)"
69
+
70
+ EXAMPLE_PROMPTS = [
71
+ ["一位穿着粉色吊带罗纹长裙的亚洲少女,外搭一件米白色毛绒短开襟衫,在阳光洒落的森林小径上侧身回眸。她拥有淡粉色薰衣草发色的甜美脸庞,发间别着一朵白色小花。黄金时段的光线穿过浓密的树叶,在深绿色的背景上形成美丽的景深光斑 和柔和光晕。电影级肖像摄影,超高画质,细腻的皮肤纹理,强调少女的温柔与唯美浪漫的日系氛围。"],
72
+ ["一幅具有电影感的胶片肖像,一位美丽的中国女生,凌乱的黑发在风中飘动遮住脸庞,眼神灵动地看着镜头。她在画面的左1/3处。她围着一条厚实的鲜红色针织围巾,穿着一件破旧的米色羊羔毛外套。背景是日落时分寒冷、干枯的荒野和远山。强烈的金色逆光直射镜头,产生巨大的镜头眩光和朦胧的光晕效果,空气中有尘埃感。胶片颗粒质感,浅景深,自然原始的风格。"],
73
+ ["一个半人半机械的黑客,坐在充满全息屏幕的黑暗房间里,绿色的代码光映照在他的脸上,赛博朋克风格,高科技细节,锐利的焦点。"],
74
+ ["A surreal double exposure portrait that blends a woman’s face with a beautiful seascape. The overall mood is dreamy and mystical, with rich colors and intricate details."],
75
+ ["A close-up, macro photography stock photo of a strawberry intricately sculpted into the shape of a hummingbird in mid-flight, its wings a blur as it sips nectar from a vibrant, tubular flower. The backdrop features a lush, colorful garden with a soft, bokeh effect, creating a dreamlike atmosphere. The image is exceptionally detailed and captured with a shallow depth of field, ensuring a razor-sharp focus on the strawberry-hummingbird and gentle fading of the background. The high resolution, professional photographers style, and soft lighting illuminate the scene in a very detailed manner, professional color grading amplifies the vibrant colors and creates an image with exceptional clarity. The depth of field makes the hummingbird and flower stand out starkly against the bokeh background."],
76
+ ["网红咖啡店内部,透过钢化玻璃拍摄,中景平视角度;玻璃表面有环境反光与色彩叠影,人物面部柔光打亮,坐着看向镜头,穿着带大毛领的宽松上衣;白天咖啡店,太阳光线打在人物脸上,玻璃反光清透自然,ccd质感。"],
77
+ ["室内中景人像摄影,复古胶片风格,电影叙事感画面。一位清纯气质的年轻女性,留着黑色齐刘海长直发,妆容清透伪素颜,皮肤白皙透亮。她身穿一件质地柔软、淡绿色的马海毛(Mohair)绒毛毛衣,质感毛绒蓬松,下身搭配淡青色棉麻长裙。人物慵懒地蜷缩/侧卧在沙发角落,身体姿态放松柔软,呈现自然的C型曲线。一只手轻轻拿着一颗鲜红的番茄靠近脸颊和下巴,眼神迷离、温柔且深情地直视镜头,表情处于放空与凝视之间,极具故事感。复古文艺的室内一角,沙发上铺着淡雅的复古碎花布艺沙发罩,身旁放着一盘红色的番茄作为前景点缀。背景虚化,隐约可见室内的陈设与绿植,整体环境色调偏向青绿色���胶片感。极具艺术感的局部自然光(丁达尔效应光斑)。一束明亮的午后阳光精准地照射在手部、手中的番茄以及面部一侧,形成强烈的明暗对比(Chiaroscuro)。高光部分带有光晕(Bloom),阴影部分呈现胶片特有的青蓝色调,光影层次丰富。。慵懒、静谧、梦幻、日系文艺、情绪感强、高级且富有夏末秋初的诗意。。模拟胶片相机(如Contax T3或Pentax 67)拍摄,使用50mm标准定焦镜头,大光圈(f/1.8)制造柔和的背景虚化。后期加入明显的粗颗粒胶片滤镜(Heavy Film Grain)和色彩偏移,增强模拟摄影的真实感与年代感。极度真实的皮肤质感,保留面部微小的毛孔和纹理,拒绝过度磨皮;马海毛毛衣在逆光下呈现出清晰的绒毛光晕边缘;番茄表面光滑的高光反射;碎花布料的褶皱细节;整体画面覆盖一层复古的胶片噪点。"],
78
+ ]
79
+
80
+ def get_resolution(resolution_str):
81
+ match = re.search(r"(\d+)\s*[×x]\s*(\d+)", resolution_str)
82
+ if match:
83
+ return int(match.group(1)), int(match.group(2))
84
+ return 1024, 1024
85
+
86
+ def load_models(model_path):
87
+ print(f"Loading BitDance model from {model_path}...")
88
+
89
+ if not os.path.exists(model_path):
90
+ print(f"Warning: Model path {model_path} does not exist locally. Attempting to load anyway (or handle download logic here).")
91
+
92
+ pipe = BitDanceT2IPipeline(model_path=model_path, device="cuda")
93
+ return pipe
94
+
95
+ def generate_image(
96
+ pipe,
97
+ prompt,
98
+ resolution,
99
+ seed=42,
100
+ guidance_scale=7.5,
101
+ num_inference_steps=50,
102
+ ):
103
+ width, height = get_resolution(resolution)
104
+
105
+ images = pipe.generate(
106
+ prompt=prompt,
107
+ height=height,
108
+ width=width,
109
+ num_sampling_steps=num_inference_steps,
110
+ guidance_scale=guidance_scale,
111
+ num_images=1,
112
+ seed=seed
113
+ )
114
+
115
+ return images[0]
116
+
117
+ pipe = None
118
+
119
+ def init_app():
120
+ global pipe
121
+ try:
122
+ pipe = load_models(MODEL_PATH)
123
+ print("Model loaded successfully.")
124
+ except Exception as e:
125
+ print(f"Error loading model: {e}")
126
+ pipe = None
127
+
128
+ def generate(
129
+ prompt,
130
+ resolution,
131
+ seed=42,
132
+ steps=50,
133
+ guidance_scale=7.5,
134
+ random_seed=True,
135
+ gallery_images=None,
136
+ progress=gr.Progress(track_tqdm=True),
137
+ ):
138
+ if random_seed:
139
+ new_seed = random.randint(1, 1000000)
140
+ else:
141
+ new_seed = seed if seed != -1 else random.randint(1, 1000000)
142
+
143
+ if pipe is None:
144
+ raise gr.Error("Model not loaded.")
145
+
146
+ print(f"Generating: Prompt='{prompt[:20]}...', Res={resolution}, Seed={new_seed}, Steps={steps}, CFG={guidance_scale}")
147
+
148
+ try:
149
+ image = generate_image(
150
+ pipe=pipe,
151
+ prompt=prompt,
152
+ resolution=resolution,
153
+ seed=new_seed,
154
+ guidance_scale=guidance_scale,
155
+ num_inference_steps=int(steps),
156
+ )
157
+ except Exception as e:
158
+ raise gr.Error(f"Generation failed: {str(e)}")
159
+
160
+ if gallery_images is None:
161
+ gallery_images = []
162
+
163
+ gallery_images = [image] + gallery_images
164
+
165
+ return gallery_images, str(new_seed), int(new_seed)
166
+
167
+ init_app()
168
+
169
+ # ==================== Gradio UI ====================
170
+
171
+ with gr.Blocks(title="BitDance Demo") as demo:
172
+ gr.Markdown(
173
+ """<div align="center">
174
+
175
+ # BitDance Generation Demo
176
+
177
+ [![GitHub](https://img.shields.io/badge/GitHub-BitDance-181717?logo=github&logoColor=white)](https://github.com/shallowdream204/BitDance)
178
+
179
+ *BitDance: Scaling Autoregressive Generative Models with Binary Tokens*
180
+
181
+ </div>"""
182
+ )
183
+
184
+ with gr.Row():
185
+ with gr.Column(scale=1):
186
+ prompt_input = gr.Textbox(label="Prompt", lines=3, placeholder="Enter your prompt here...")
187
+
188
+ resolution = gr.Dropdown(
189
+ value=DEFAULT_RES,
190
+ choices=RESOLUTION_CHOICES,
191
+ label="Resolution (Width x Height)"
192
+ )
193
+
194
+ with gr.Row():
195
+ seed = gr.Number(label="Seed", value=42, precision=0)
196
+ random_seed = gr.Checkbox(label="Random Seed", value=True)
197
+
198
+ with gr.Row():
199
+ steps = gr.Slider(label="Diffusion Sampling Steps", minimum=10, maximum=100, value=50, step=1)
200
+ guidance_scale = gr.Slider(label="CFG Guidance Scale", minimum=1.0, maximum=15.0, value=7.5, step=0.5)
201
+
202
+ generate_btn = gr.Button("Generate", variant="primary")
203
+
204
+ gr.Markdown("### 📝 Example Prompts")
205
+ gr.Examples(examples=EXAMPLE_PROMPTS, inputs=prompt_input, label=None)
206
+
207
+ with gr.Column(scale=1):
208
+ output_gallery = gr.Gallery(
209
+ label="Generated Images",
210
+ columns=2,
211
+ rows=2,
212
+ height=600,
213
+ object_fit="contain",
214
+ format="png",
215
+ interactive=False,
216
+ )
217
+ used_seed = gr.Textbox(label="Seed Used", interactive=False)
218
+
219
+ generate_btn.click(
220
+ generate,
221
+ inputs=[prompt_input, resolution, seed, steps, guidance_scale, random_seed, output_gallery],
222
+ outputs=[output_gallery, used_seed, seed],
223
+ api_visibility="public",
224
+ )
225
+
226
+ css = """
227
+ .fillable{max-width: 1230px !important}
228
+ """
229
+
230
+ if __name__ == "__main__":
231
+ demo.launch(css=css, mcp_server=True)
modeling/t2i_pipeline.py ADDED
@@ -0,0 +1,283 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+ from torch import nn
3
+ from einops import rearrange
4
+
5
+ from transformers import set_seed
6
+ from PIL import Image
7
+
8
+ import numpy as np
9
+ from torch import nn
10
+ from transformers import AutoTokenizer, Qwen3ForCausalLM, Qwen3Config
11
+
12
+ from modeling.utils import MLPconnector
13
+ from modeling.vision_encoder.autoencoder import VQModel
14
+ from modeling.vision_head.flow_head_parallel_x import DiffHead
15
+
16
+ from safetensors.torch import load_file as load_sft
17
+ import json
18
+ import os
19
+ from tqdm import tqdm
20
+
21
+ IMAGE_SIZE_LIST = [
22
+ # --- 1024px Area ---
23
+ [2048, 512],
24
+ [1920, 512],
25
+ [1536, 640],
26
+ [1280, 768],
27
+ [1152, 896],
28
+ [1024, 1024],
29
+ [896, 1152],
30
+ [768, 1280],
31
+ [640, 1536],
32
+ [512, 1920],
33
+ [512, 2048],
34
+ # --- 512px Area ---
35
+ [1024, 256],
36
+ [896, 256],
37
+ [640, 384],
38
+ [512, 512],
39
+ [384, 640],
40
+ [256, 896],
41
+ [256, 1024],
42
+ ]
43
+
44
+ class BitDanceT2IPipeline:
45
+ def __init__(self, model_path, device='cuda'):
46
+ self.device = device
47
+ # LLM and tokenizer
48
+ self.tokenizer = AutoTokenizer.from_pretrained(model_path)
49
+ self.llm_config = Qwen3Config.from_pretrained(model_path)
50
+ self.llm_model = Qwen3ForCausalLM.from_pretrained(model_path, torch_dtype=torch.bfloat16).eval().to(device)
51
+ self.hidden_size = self.llm_config.hidden_size
52
+
53
+ # Autoencoder
54
+ with open(os.path.join(model_path, 'ae_config.json'), "r") as f:
55
+ self.ae_config = json.load(f)
56
+ self.ae = VQModel(**self.ae_config).eval()
57
+ self.ae.load_state_dict(load_sft(os.path.join(model_path, 'ae.safetensors')), strict=True, assign=True)
58
+ self.ae.to(device)
59
+ self.vae_patch_size = 2 ** (len(self.ae_config['ddconfig']['ch_mult'])-1)
60
+
61
+ # Vision head
62
+ with open(os.path.join(model_path, 'vision_head_config.json'), "r") as f:
63
+ self.vision_head_config = json.load(f)
64
+ self.vision_head = DiffHead(**self.vision_head_config).eval()
65
+ self.vision_head.load_state_dict(load_sft(os.path.join(model_path, 'vision_head.safetensors')), strict=True, assign=True)
66
+ self.vision_head.to(device)
67
+ self.parallel_num = self.vision_head_config['parallel_num']
68
+ print(f'use {self.parallel_num}-token parallel prediction per step')
69
+ self.ps = int(self.parallel_num ** 0.5)
70
+
71
+ # Projector
72
+ self.embed_vision_mlp = MLPconnector(self.ae_config['ddconfig']['z_channels'], self.hidden_size, "gelu_pytorch_tanh")
73
+ self.embed_vision_mlp.load_state_dict(load_sft(os.path.join(model_path, 'projector.safetensors')), strict=True, assign=True)
74
+ self.embed_vision_mlp.to(device)
75
+
76
+ # 2D sinusoidal position embedding
77
+ self.build_pos_embed()
78
+
79
+ def build_pos_embed(self, max_len=4096):
80
+ max_len = max_len // self.vae_patch_size
81
+ pos_embed_1d = self._get_1d_sincos_pos_embed(self.hidden_size//2, max_len)
82
+ pos_embed_1d = nn.Parameter(pos_embed_1d, requires_grad=False)
83
+ self.pos_embed_1d = pos_embed_1d.to(self.device)
84
+
85
+ def _get_1d_sincos_pos_embed(self, dim, max_len, pe_interpolation=1.0):
86
+ assert dim % 2 == 0
87
+ omega = torch.arange(dim // 2, dtype=torch.float32)
88
+ omega /= dim / 2.0
89
+ omega = 1.0 / 10000**omega # (D/4,)
90
+
91
+ pos = torch.arange(max_len, dtype=torch.float32) / pe_interpolation
92
+ out = torch.einsum("m,d->md", pos, omega) # (max_len, D/4)
93
+
94
+ emb_sin = torch.sin(out)
95
+ emb_cos = torch.cos(out)
96
+ return torch.cat([emb_sin, emb_cos], dim=1) # (max_len, D/2)
97
+
98
+ def get_2d_embed(self, h, w, ps=1):
99
+ emb_v = self.pos_embed_1d[:h]
100
+ emb_h = self.pos_embed_1d[:w]
101
+
102
+ grid_v = emb_v.view(h, 1, self.hidden_size//2).repeat(1, w, 1)
103
+ grid_h = emb_h.view(1, w, self.hidden_size//2).repeat(h, 1, 1)
104
+
105
+ pos_embed = torch.cat([grid_h, grid_v], dim=-1) # h w c
106
+
107
+ return rearrange(pos_embed, '(h p1) (w p2) c -> (h w p1 p2) c', p1=ps, p2=ps)
108
+
109
+ @torch.no_grad()
110
+ def generate(
111
+ self,
112
+ prompt: str,
113
+ height: int = 1024,
114
+ width: int = 1024,
115
+ num_sampling_steps: int = 50,
116
+ guidance_scale: float = 7.5,
117
+ num_images: int = 1,
118
+ seed: int = 1234,
119
+ ):
120
+ # Set seed for reproducibility
121
+ if seed is not None:
122
+ set_seed(seed)
123
+ # Calculate max_length dynamically based on image_size and stride of 16
124
+ max_length = (height // self.vae_patch_size) * (width // self.vae_patch_size)
125
+
126
+ image_size = [height, width]
127
+ if image_size not in IMAGE_SIZE_LIST:
128
+ raise ValueError(f"image_size {image_size} is not supported. Please choose from {IMAGE_SIZE_LIST}")
129
+
130
+ with torch.amp.autocast("cuda", enabled=True, dtype=torch.bfloat16):
131
+ gen_images = self.gen_image(
132
+ cond_prompt=f"<|im_start|>user\n{prompt}<|im_end|>\n<|im_start|>assistant\n",
133
+ uncond_prompt="<|im_start|>assistant\n",
134
+ guidance_scale=guidance_scale,
135
+ num_sampling_steps=num_sampling_steps,
136
+ num_images=num_images,
137
+ image_size=image_size,
138
+ max_length=max_length,
139
+ show_progress=True,
140
+ )
141
+
142
+ gen_images = (
143
+ torch.clamp(127.5 * gen_images + 128.0, 0, 255)
144
+ .permute(0, 2, 3, 1)
145
+ .to("cpu", dtype=torch.uint8)
146
+ .numpy()
147
+ )
148
+ pil_images = []
149
+ for i in range(gen_images.shape[0]):
150
+ img_array = gen_images[i]
151
+ if img_array.dtype != np.uint8:
152
+ img_array = img_array.astype(np.uint8)
153
+ pil_images.append(Image.fromarray(img_array))
154
+
155
+ return pil_images
156
+
157
+ @torch.no_grad()
158
+ def gen_image(self,
159
+ cond_prompt,
160
+ uncond_prompt=None,
161
+ guidance_scale: float = 1.0,
162
+ num_sampling_steps: int = 50,
163
+ max_length: int = 64,
164
+ num_images: int = 1,
165
+ image_size = [256, 256],
166
+ show_progress: bool = False,
167
+ ):
168
+ tokenizer = self.tokenizer
169
+ device = self.device
170
+ model = self.llm_model.model
171
+
172
+ step_width = self.parallel_num
173
+ num_steps = max_length // step_width
174
+
175
+ cond_ids = torch.tensor(tokenizer.encode(cond_prompt), device=device, dtype=torch.long)
176
+ cond_emb = model.embed_tokens(cond_ids)
177
+ if guidance_scale > 1.0:
178
+ uncond_ids = torch.tensor(tokenizer.encode(uncond_prompt), device=device, dtype=torch.long)
179
+ uncond_emb = model.embed_tokens(uncond_ids)
180
+
181
+ img_start_id = tokenizer.convert_tokens_to_ids("<|vision_start|>")
182
+ res_h_token_id = tokenizer.convert_tokens_to_ids(f"<|res_{image_size[0] // self.vae_patch_size}|>")
183
+ res_w_token_id = tokenizer.convert_tokens_to_ids(f"<|res_{image_size[1] // self.vae_patch_size}|>")
184
+ img_start_emb = model.embed_tokens(torch.tensor([img_start_id, res_h_token_id, res_w_token_id], device=device))
185
+
186
+ h, w = image_size[0] // self.vae_patch_size, image_size[1] // self.vae_patch_size
187
+ # prepare diff pos embed
188
+ pos_embed_for_diff = self.get_2d_embed(h, w, ps=self.ps if hasattr(self, 'ps') else 1).unsqueeze(0)
189
+
190
+ # add query tokens for parallel decoding
191
+ for i in range(1, self.parallel_num):
192
+ query_token = torch.tensor([tokenizer.convert_tokens_to_ids(f"<|query_{i}|>")], device=self.device, dtype=torch.long)
193
+ query_embed = self.llm_model.model.embed_tokens(query_token)
194
+ img_start_emb = torch.cat([img_start_emb, query_embed], dim=0)
195
+
196
+ input_embeds_cond = torch.cat(
197
+ [cond_emb, img_start_emb], dim=0
198
+ ).unsqueeze(0).repeat(num_images, 1, 1)
199
+ outputs_c = model(
200
+ inputs_embeds=input_embeds_cond[:, :-step_width, :],
201
+ use_cache=True,
202
+ )
203
+ pkv_c = outputs_c.past_key_values
204
+
205
+ # bidirectional attn
206
+ bi_attn_mask = torch.ones(
207
+ (input_embeds_cond.shape[0], 1, step_width, step_width+pkv_c[0][0].shape[2]),
208
+ dtype=torch.bool,
209
+ device=device,
210
+ )
211
+ outputs_c = model(
212
+ inputs_embeds=input_embeds_cond[:, -step_width:, :],
213
+ past_key_values=pkv_c,
214
+ use_cache=True,
215
+ attention_mask=bi_attn_mask,
216
+ )
217
+ pkv_c = outputs_c.past_key_values
218
+ hidden_c = outputs_c.last_hidden_state[:, -step_width:] # [B, parallel_num, D]
219
+
220
+ if guidance_scale > 1.0:
221
+ input_embeds_uncond = torch.cat(
222
+ [uncond_emb, img_start_emb], dim=0
223
+ ).unsqueeze(0).repeat(num_images, 1, 1)
224
+ outputs_u = model(
225
+ inputs_embeds=input_embeds_uncond[:, :-step_width, :],
226
+ use_cache=True,
227
+ )
228
+ pkv_u = outputs_u.past_key_values
229
+ outputs_u = model(
230
+ inputs_embeds=input_embeds_uncond[:, -step_width:, :],
231
+ past_key_values=pkv_u,
232
+ use_cache=True,
233
+ attention_mask=bi_attn_mask,
234
+ )
235
+ pkv_u = outputs_u.past_key_values
236
+ hidden_u = outputs_u.last_hidden_state[:, -step_width:] # [B, parallel_num, D]
237
+
238
+ out_tokens = []
239
+ if show_progress:
240
+ pbar = tqdm(total=num_steps, desc="Decoding Steps")
241
+ for step in range(num_steps):
242
+ if show_progress:
243
+ pbar.update(1)
244
+ h_fused = torch.cat([hidden_c, hidden_u], dim=0) if guidance_scale > 1.0 else hidden_c
245
+ h_fused = h_fused + pos_embed_for_diff[:, step*step_width:(step+1)*step_width, :]
246
+ pred_latents = self.vision_head.sample(h_fused, num_sampling_steps=num_sampling_steps, cfg=guidance_scale)
247
+ # important! LFQ is used here
248
+ curr_tokens = torch.sign(pred_latents)
249
+ curr_embeds = self.embed_vision_mlp(curr_tokens)
250
+ out_tokens.append(curr_tokens[:num_images])
251
+ model_input = curr_embeds # [B, N, D]
252
+ # 2d pos embed
253
+ model_input = model_input + pos_embed_for_diff[:, step*step_width:(step+1)*step_width, :]
254
+
255
+ # bidirectional attn mask
256
+ bi_attn_mask = torch.ones(
257
+ (model_input.shape[0], 1, model_input.shape[1], model_input.shape[1]+pkv_c[0][0].shape[2]),
258
+ dtype=torch.bool,
259
+ device=device
260
+ )
261
+ outputs_c = model(inputs_embeds=model_input[:num_images], past_key_values=pkv_c, use_cache=True, attention_mask=bi_attn_mask[:num_images])
262
+ pkv_c = outputs_c.past_key_values
263
+ hidden_c = outputs_c.last_hidden_state[:, -step_width:] # [B, parallel_num, D]
264
+
265
+ if guidance_scale > 1.0:
266
+ outputs_u = model(inputs_embeds=model_input[num_images:], past_key_values=pkv_u, use_cache=True, attention_mask=bi_attn_mask[num_images:])
267
+ pkv_u = outputs_u.past_key_values
268
+ hidden_u = outputs_u.last_hidden_state[:, -step_width:] # [B, parallel_num, D]
269
+
270
+ full_output = torch.cat(out_tokens, dim=1)
271
+ image = self.decode_image(full_output, [h, w], ps=self.ps if hasattr(self, 'ps') else 1) # [num_images, c, h, w]
272
+ return image
273
+
274
+ def decode_image(self, image_latents, image_size=None, ps=1):
275
+ if image_size is None:
276
+ h = w = int(image_latents.size(1) ** 0.5)
277
+ else:
278
+ h, w = image_size
279
+
280
+ image_latents = rearrange(image_latents, 'b (h w p1 p2) c -> b c (h p1) (w p2)', h=h//ps, w=w//ps, p1=ps, p2=ps)
281
+ output = self.ae.decode(image_latents) # [1, c, h, w]
282
+
283
+ return output
modeling/utils.py ADDED
@@ -0,0 +1,216 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+ import torch.nn.functional as F
3
+ from torch import nn
4
+
5
+ from torch.nn.attention.flex_attention import or_masks, and_masks
6
+
7
+ from transformers.activations import ACT2FN
8
+
9
+ class MLPconnector(nn.Module):
10
+ def __init__(self, in_dim: int, out_dim: int, hidden_act: str):
11
+ super().__init__()
12
+ self.activation_fn = ACT2FN[hidden_act]
13
+ self.fc1 = nn.Linear(in_dim, out_dim)
14
+ self.fc2 = nn.Linear(out_dim, out_dim)
15
+
16
+ def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
17
+ hidden_states = self.fc1(hidden_states)
18
+ hidden_states = self.activation_fn(hidden_states)
19
+ hidden_states = self.fc2(hidden_states)
20
+ return hidden_states
21
+
22
+ def create_sparse_mask(document_lens, split_lens, attn_modes, parallel_num, device):
23
+ parallel_causal_num = 2
24
+ parallel_block_causal_num = parallel_num
25
+
26
+ def causal_mask(b, h, q_idx, kv_idx):
27
+ return q_idx >= kv_idx
28
+
29
+ def parallel_block_mask(b, h, q_idx, kv_idx):
30
+ same_seg = segment_ids[q_idx] == segment_ids[kv_idx]
31
+ is_par = is_parallel[q_idx]
32
+
33
+ lq = local_ids[q_idx]
34
+ lk = local_ids[kv_idx]
35
+
36
+ in_block_region = (lq >= parallel_causal_num) & (lk >= parallel_causal_num)
37
+
38
+ same_block = ((lq - parallel_causal_num) // parallel_block_causal_num) == ((lk - parallel_causal_num) // parallel_block_causal_num)
39
+
40
+ return same_seg & is_par & in_block_region & same_block
41
+
42
+ def sample_mask(b, h, q_idx, kv_idx):
43
+ return document_id[q_idx] == document_id[kv_idx]
44
+
45
+ segment_ids_list = []
46
+ local_ids_list = []
47
+ is_parallel_list = []
48
+
49
+ current_seg_id = 0
50
+ for length, mode in zip(split_lens, attn_modes):
51
+ segment_ids_list.extend([current_seg_id] * length)
52
+ local_ids_list.extend(list(range(length)))
53
+ is_parallel_list.extend([True if mode == 'parallel' else False] * length)
54
+ current_seg_id += 1
55
+
56
+ segment_ids = torch.tensor(segment_ids_list, device=device, dtype=torch.long)
57
+ local_ids = torch.tensor(local_ids_list, device=device, dtype=torch.long)
58
+ is_parallel = torch.tensor(is_parallel_list, device=device, dtype=torch.bool)
59
+
60
+ document_id = torch.cat([torch.full((l,), i, device=device) for i, l in enumerate(document_lens, start=1)])
61
+
62
+ return and_masks(or_masks(causal_mask, parallel_block_mask), sample_mask)
63
+
64
+ def top_k_top_p_filtering(
65
+ logits,
66
+ top_k: int = 0,
67
+ top_p: float = 1.0,
68
+ filter_value: float = -float("Inf"),
69
+ min_tokens_to_keep: int = 1,
70
+ ):
71
+ """Filter a distribution of logits using top-k and/or top-p (nucleus) filtering."""
72
+ if top_k > 0:
73
+ top_k = min(max(top_k, min_tokens_to_keep), logits.size(-1))
74
+ indices_to_remove = logits < torch.topk(logits, top_k)[0][..., -1, None]
75
+ logits[indices_to_remove] = filter_value
76
+
77
+ if top_p < 1.0:
78
+ sorted_logits, sorted_indices = torch.sort(logits, descending=True)
79
+ cumulative_probs = torch.cumsum(F.softmax(sorted_logits, dim=-1), dim=-1)
80
+
81
+ sorted_indices_to_remove = cumulative_probs > top_p
82
+ if min_tokens_to_keep > 1:
83
+ sorted_indices_to_remove[..., :min_tokens_to_keep] = 0
84
+ sorted_indices_to_remove[..., 1:] = sorted_indices_to_remove[..., :-1].clone()
85
+ sorted_indices_to_remove[..., 0] = 0
86
+
87
+ indices_to_remove = sorted_indices_to_remove.scatter(1, sorted_indices, sorted_indices_to_remove)
88
+ logits[indices_to_remove] = filter_value
89
+
90
+ return logits
91
+
92
+
93
+ def sample_codebook(
94
+ pred_logits,
95
+ cur_item_type,
96
+ codebook,
97
+ do_sample: bool = True,
98
+ temperature: float = 1.0,
99
+ top_k: int = 0,
100
+ top_p: float = 1.0,
101
+ ):
102
+ """
103
+ pred_logits: (B, vocab_size)
104
+ cur_item_type: 'text' or 'vision'
105
+ """
106
+ # 1. Apply temperature
107
+ logits = pred_logits / max(temperature, 1e-5)
108
+
109
+ # 2. Apply top-k / top-p filtering
110
+ if top_k > 0 or top_p < 1.0:
111
+ logits = top_k_top_p_filtering(logits, top_k=top_k, top_p=top_p)
112
+
113
+ # 3. Get probabilities
114
+ probs = F.softmax(logits, dim=-1)
115
+
116
+ # 4. Sample or take argmax
117
+ if do_sample:
118
+ curr_tokens = torch.multinomial(probs, num_samples=1).squeeze(-1)
119
+ else:
120
+ curr_tokens = torch.argmax(probs, dim=-1)
121
+
122
+ curr_embeds = codebook(curr_tokens)
123
+
124
+ return curr_tokens, curr_embeds
125
+
126
+
127
+ def flip_tensor_elements_uniform_prob(tensor: torch.Tensor, p_max: float) -> torch.Tensor:
128
+ if not 0.0 <= p_max <= 1.0:
129
+ raise ValueError(f"p_max must in [0.0, 1.0]")
130
+
131
+ r1 = torch.rand_like(tensor)
132
+ r2 = torch.rand_like(tensor)
133
+
134
+ flip_mask = r1 < p_max * r2
135
+
136
+ multiplier = torch.where(flip_mask, -1.0, 1.0)
137
+ multiplier = multiplier.to(tensor.dtype)
138
+
139
+ flipped_tensor = tensor * multiplier
140
+ return flipped_tensor
141
+
142
+ def gaussian_sample(raw_output):
143
+ mu, log_var = raw_output.chunk(2, dim=-1)
144
+ sigma = torch.exp(0.5 * log_var)
145
+ sample = mu + torch.randn_like(mu) * sigma
146
+
147
+ return sample
148
+
149
+ def get_2d_sincos_pos_embed(embed_dim, grid_size, cls_token=False, extra_tokens=0, pe_interpolation=1.0):
150
+ """
151
+ grid_size: int or tuple/list of (h, w)
152
+ return:
153
+ pos_embed: [grid_h*grid_w, embed_dim] or [extra_tokens+grid_h*grid_w, embed_dim] (w/ or w/o cls_token)
154
+ """
155
+ if isinstance(grid_size, int):
156
+ grid_h_size, grid_w_size = grid_size, grid_size
157
+ else:
158
+ grid_h_size, grid_w_size = grid_size
159
+
160
+ grid_h = torch.arange(grid_h_size, dtype=torch.float32) / pe_interpolation
161
+ grid_w = torch.arange(grid_w_size, dtype=torch.float32) / pe_interpolation
162
+
163
+ grid_w, grid_h = torch.meshgrid(grid_w, grid_h, indexing='xy')
164
+
165
+ grid = torch.stack([grid_w, grid_h], dim=0) # shape: (2, grid_h_size, grid_w_size)
166
+
167
+ grid = grid.reshape([2, 1, grid_h_size, grid_w_size])
168
+ pos_embed = get_2d_sincos_pos_embed_from_grid(embed_dim, grid)
169
+
170
+ if cls_token and extra_tokens > 0:
171
+ pos_embed = torch.cat([torch.zeros([extra_tokens, embed_dim]), pos_embed], dim=0)
172
+
173
+ return pos_embed
174
+
175
+ def get_2d_sincos_pos_embed_from_grid(embed_dim, grid):
176
+ assert embed_dim % 2 == 0
177
+
178
+ # use half of dimensions to encode
179
+ emb_h = get_1d_sincos_pos_embed_from_grid(embed_dim // 2, grid[0]) # (H*W, D/2)
180
+ emb_w = get_1d_sincos_pos_embed_from_grid(embed_dim // 2, grid[1]) # (H*W, D/2)
181
+
182
+ emb = torch.cat([emb_h, emb_w], dim=1) # (H*W, D)
183
+ return emb
184
+
185
+ def get_1d_sincos_pos_embed_from_grid(embed_dim, pos):
186
+ """
187
+ embed_dim: output dimension for each position
188
+ pos: a list of positions to be encoded: size (M,)
189
+ out: (M, D)
190
+ """
191
+ assert embed_dim % 2 == 0
192
+ omega = torch.arange(embed_dim // 2, dtype=torch.float32)
193
+ omega /= embed_dim / 2.0
194
+ omega = 1.0 / 10000**omega # (D/2,)
195
+
196
+ pos = pos.reshape(-1) # (M,)
197
+
198
+ out = torch.einsum("m,d->md", pos, omega) # (M, D/2), outer product
199
+
200
+ emb_sin = torch.sin(out) # (M, D/2)
201
+ emb_cos = torch.cos(out) # (M, D/2)
202
+
203
+ emb = torch.cat([emb_sin, emb_cos], dim=1) # (M, D)
204
+ return emb
205
+
206
+ def remove_first_user_block(x: str) -> str:
207
+ start_marker = "<|im_start|>user\n"
208
+ end_marker = "<|im_end|>\n"
209
+ start_index = x.find(start_marker)
210
+ if start_index == -1:
211
+ return x
212
+ end_index = x.find(end_marker, start_index + len(start_marker))
213
+ if end_index == -1:
214
+ return x
215
+ result = x[:start_index] + x[end_index + len(end_marker):]
216
+ return result
modeling/vision_encoder/autoencoder.py ADDED
@@ -0,0 +1,520 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+ import torch.nn as nn
3
+ import math
4
+ from einops import rearrange
5
+ import torch.nn.functional as F
6
+ from collections import defaultdict
7
+
8
+
9
+ def swish(x):
10
+ return x*torch.sigmoid(x)
11
+
12
+ class ResBlock(nn.Module):
13
+ def __init__(self,
14
+ in_filters,
15
+ out_filters,
16
+ use_conv_shortcut = False,
17
+ use_agn = False,
18
+ ) -> None:
19
+ super().__init__()
20
+
21
+ self.in_filters = in_filters
22
+ self.out_filters = out_filters
23
+ self.use_conv_shortcut = use_conv_shortcut
24
+ self.use_agn = use_agn
25
+
26
+ if not use_agn: ## agn is GroupNorm likewise skip it if has agn before
27
+ self.norm1 = nn.GroupNorm(32, in_filters, eps=1e-6)
28
+ self.norm2 = nn.GroupNorm(32, out_filters, eps=1e-6)
29
+
30
+ self.conv1 = nn.Conv2d(in_filters, out_filters, kernel_size=(3, 3), padding=1, bias=False)
31
+ self.conv2 = nn.Conv2d(out_filters, out_filters, kernel_size=(3, 3), padding=1, bias=False)
32
+
33
+ if in_filters != out_filters:
34
+ if self.use_conv_shortcut:
35
+ self.conv_shortcut = nn.Conv2d(in_filters, out_filters, kernel_size=(3, 3), padding=1, bias=False)
36
+ else:
37
+ self.nin_shortcut = nn.Conv2d(in_filters, out_filters, kernel_size=(1, 1), padding=0, bias=False)
38
+
39
+
40
+ def forward(self, x, **kwargs):
41
+ residual = x
42
+
43
+ if not self.use_agn:
44
+ x = self.norm1(x)
45
+ x = swish(x)
46
+ x = self.conv1(x)
47
+ x = self.norm2(x)
48
+ x = swish(x)
49
+ x = self.conv2(x)
50
+ if self.in_filters != self.out_filters:
51
+ if self.use_conv_shortcut:
52
+ residual = self.conv_shortcut(residual)
53
+ else:
54
+ residual = self.nin_shortcut(residual)
55
+
56
+ return x + residual
57
+
58
+ class Encoder(nn.Module):
59
+ def __init__(self, *, ch, out_ch, in_channels, num_res_blocks, z_channels, ch_mult=(1, 2, 2, 4),
60
+ resolution=None, double_z=False,
61
+ ):
62
+ super().__init__()
63
+
64
+ self.in_channels = in_channels
65
+ self.z_channels = z_channels
66
+ self.resolution = resolution
67
+
68
+ self.num_res_blocks = num_res_blocks
69
+ self.num_blocks = len(ch_mult)
70
+
71
+ self.conv_in = nn.Conv2d(in_channels,
72
+ ch,
73
+ kernel_size=(3, 3),
74
+ padding=1,
75
+ bias=False
76
+ )
77
+
78
+ ## construct the model
79
+ self.down = nn.ModuleList()
80
+
81
+ in_ch_mult = (1,)+tuple(ch_mult)
82
+ for i_level in range(self.num_blocks):
83
+ block = nn.ModuleList()
84
+ block_in = ch*in_ch_mult[i_level] #[1, 1, 2, 2, 4]
85
+ block_out = ch*ch_mult[i_level] #[1, 2, 2, 4]
86
+ for _ in range(self.num_res_blocks):
87
+ block.append(ResBlock(block_in, block_out))
88
+ block_in = block_out
89
+
90
+ down = nn.Module()
91
+ down.block = block
92
+ if i_level < self.num_blocks - 1:
93
+ down.downsample = nn.Conv2d(block_out, block_out, kernel_size=(3, 3), stride=(2, 2), padding=1)
94
+
95
+ self.down.append(down)
96
+
97
+ ### mid
98
+ self.mid_block = nn.ModuleList()
99
+ for res_idx in range(self.num_res_blocks):
100
+ self.mid_block.append(ResBlock(block_in, block_in))
101
+
102
+ ### end
103
+ self.norm_out = nn.GroupNorm(32, block_out, eps=1e-6)
104
+ self.conv_out = nn.Conv2d(block_out, z_channels, kernel_size=(1, 1))
105
+
106
+ def forward(self, x):
107
+
108
+ ## down
109
+ x = self.conv_in(x)
110
+ for i_level in range(self.num_blocks):
111
+ for i_block in range(self.num_res_blocks):
112
+ x = self.down[i_level].block[i_block](x)
113
+
114
+ if i_level < self.num_blocks - 1:
115
+ x = self.down[i_level].downsample(x)
116
+
117
+ ## mid
118
+ for res in range(self.num_res_blocks):
119
+ x = self.mid_block[res](x)
120
+
121
+
122
+ x = self.norm_out(x)
123
+ x = swish(x)
124
+ x = self.conv_out(x)
125
+
126
+ return x
127
+
128
+ class Decoder(nn.Module):
129
+ def __init__(self, *, ch, out_ch, in_channels, num_res_blocks, z_channels, ch_mult=(1, 2, 2, 4),
130
+ resolution=None, double_z=False,) -> None:
131
+ super().__init__()
132
+
133
+ self.ch = ch
134
+ self.num_blocks = len(ch_mult)
135
+ self.num_res_blocks = num_res_blocks
136
+ self.resolution = resolution
137
+ self.in_channels = in_channels
138
+
139
+ block_in = ch*ch_mult[self.num_blocks-1]
140
+
141
+ self.conv_in = nn.Conv2d(
142
+ z_channels, block_in, kernel_size=(3, 3), padding=1, bias=True
143
+ )
144
+
145
+ self.mid_block = nn.ModuleList()
146
+ for res_idx in range(self.num_res_blocks):
147
+ self.mid_block.append(ResBlock(block_in, block_in))
148
+
149
+ self.up = nn.ModuleList()
150
+
151
+ self.adaptive = nn.ModuleList()
152
+
153
+ for i_level in reversed(range(self.num_blocks)):
154
+ block = nn.ModuleList()
155
+ block_out = ch*ch_mult[i_level]
156
+ self.adaptive.insert(0, AdaptiveGroupNorm(z_channels, block_in))
157
+ for i_block in range(self.num_res_blocks):
158
+ block.append(ResBlock(block_in, block_out))
159
+ block_in = block_out
160
+
161
+ up = nn.Module()
162
+ up.block = block
163
+ if i_level > 0:
164
+ up.upsample = Upsampler(block_in)
165
+ self.up.insert(0, up)
166
+
167
+ self.norm_out = nn.GroupNorm(32, block_in, eps=1e-6)
168
+
169
+ self.conv_out = nn.Conv2d(block_in, out_ch, kernel_size=(3, 3), padding=1)
170
+
171
+ def forward(self, z):
172
+
173
+ style = z.clone() #for adaptive groupnorm
174
+
175
+ z = self.conv_in(z)
176
+
177
+ ## mid
178
+ for res in range(self.num_res_blocks):
179
+ z = self.mid_block[res](z)
180
+
181
+ ## upsample
182
+ for i_level in reversed(range(self.num_blocks)):
183
+ ### pass in each resblock first adaGN
184
+ z = self.adaptive[i_level](z, style)
185
+ for i_block in range(self.num_res_blocks):
186
+ z = self.up[i_level].block[i_block](z)
187
+
188
+ if i_level > 0:
189
+ z = self.up[i_level].upsample(z)
190
+
191
+ z = self.norm_out(z)
192
+ z = swish(z)
193
+ z = self.conv_out(z)
194
+
195
+ return z
196
+
197
+ def depth_to_space(x: torch.Tensor, block_size: int) -> torch.Tensor:
198
+ """ Depth-to-Space DCR mode (depth-column-row) core implementation.
199
+
200
+ Args:
201
+ x (torch.Tensor): input tensor. The channels-first (*CHW) layout is supported.
202
+ block_size (int): block side size
203
+ """
204
+ # check inputs
205
+ if x.dim() < 3:
206
+ raise ValueError(
207
+ f"Expecting a channels-first (*CHW) tensor of at least 3 dimensions"
208
+ )
209
+ c, h, w = x.shape[-3:]
210
+
211
+ s = block_size**2
212
+ if c % s != 0:
213
+ raise ValueError(
214
+ f"Expecting a channels-first (*CHW) tensor with C divisible by {s}, but got C={c} channels"
215
+ )
216
+
217
+ outer_dims = x.shape[:-3]
218
+
219
+ # splitting two additional dimensions from the channel dimension
220
+ x = x.view(-1, block_size, block_size, c // s, h, w)
221
+
222
+ # putting the two new dimensions along H and W
223
+ x = x.permute(0, 3, 4, 1, 5, 2)
224
+
225
+ # merging the two new dimensions with H and W
226
+ x = x.contiguous().view(*outer_dims, c // s, h * block_size,
227
+ w * block_size)
228
+
229
+ return x
230
+
231
+ class Upsampler(nn.Module):
232
+ def __init__(
233
+ self,
234
+ dim,
235
+ dim_out = None
236
+ ):
237
+ super().__init__()
238
+ dim_out = dim * 4
239
+ self.conv1 = nn.Conv2d(dim, dim_out, (3, 3), padding=1)
240
+ self.depth2space = depth_to_space
241
+
242
+ def forward(self, x):
243
+ """
244
+ input_image: [B C H W]
245
+ """
246
+ out = self.conv1(x)
247
+ out = self.depth2space(out, block_size=2)
248
+ return out
249
+
250
+ class AdaptiveGroupNorm(nn.Module):
251
+ def __init__(self, z_channel, in_filters, num_groups=32, eps=1e-6):
252
+ super().__init__()
253
+ self.gn = nn.GroupNorm(num_groups=32, num_channels=in_filters, eps=eps, affine=False)
254
+ # self.lin = nn.Linear(z_channels, in_filters * 2)
255
+ self.gamma = nn.Linear(z_channel, in_filters)
256
+ self.beta = nn.Linear(z_channel, in_filters)
257
+ self.eps = eps
258
+
259
+ def forward(self, x, quantizer):
260
+ B, C, _, _ = x.shape
261
+ # quantizer = F.adaptive_avg_pool2d(quantizer, (1, 1))
262
+ ### calcuate var for scale
263
+ scale = rearrange(quantizer, "b c h w -> b c (h w)")
264
+ scale = scale.var(dim=-1) + self.eps #not unbias
265
+ scale = scale.sqrt()
266
+ scale = self.gamma(scale).view(B, C, 1, 1)
267
+
268
+ ### calculate mean for bias
269
+ bias = rearrange(quantizer, "b c h w -> b c (h w)")
270
+ bias = bias.mean(dim=-1)
271
+ bias = self.beta(bias).view(B, C, 1, 1)
272
+
273
+ x = self.gn(x)
274
+ x = scale * x + bias
275
+
276
+ return x
277
+
278
+ class GANDecoder(nn.Module):
279
+ def __init__(self, *, ch, out_ch, in_channels, num_res_blocks, z_channels, ch_mult=(1, 2, 2, 4),
280
+ resolution=None, double_z=False,) -> None:
281
+ super().__init__()
282
+
283
+ self.ch = ch
284
+ self.num_blocks = len(ch_mult)
285
+ self.num_res_blocks = num_res_blocks
286
+ self.resolution = resolution
287
+ self.in_channels = in_channels
288
+
289
+ block_in = ch*ch_mult[self.num_blocks-1]
290
+
291
+ self.conv_in = nn.Conv2d(
292
+ z_channels * 2, block_in, kernel_size=(3, 3), padding=1, bias=True
293
+ )
294
+
295
+ self.mid_block = nn.ModuleList()
296
+ for res_idx in range(self.num_res_blocks):
297
+ self.mid_block.append(ResBlock(block_in, block_in))
298
+
299
+ self.up = nn.ModuleList()
300
+
301
+ self.adaptive = nn.ModuleList()
302
+
303
+ for i_level in reversed(range(self.num_blocks)):
304
+ block = nn.ModuleList()
305
+ block_out = ch*ch_mult[i_level]
306
+ self.adaptive.insert(0, AdaptiveGroupNorm(z_channels, block_in))
307
+ for i_block in range(self.num_res_blocks):
308
+ # if i_block == 0:
309
+ # block.append(ResBlock(block_in, block_out, use_agn=True))
310
+ # else:
311
+ block.append(ResBlock(block_in, block_out))
312
+ block_in = block_out
313
+
314
+ up = nn.Module()
315
+ up.block = block
316
+ if i_level > 0:
317
+ up.upsample = Upsampler(block_in)
318
+ self.up.insert(0, up)
319
+
320
+ self.norm_out = nn.GroupNorm(32, block_in, eps=1e-6)
321
+
322
+ self.conv_out = nn.Conv2d(block_in, out_ch, kernel_size=(3, 3), padding=1)
323
+
324
+ def forward(self, z):
325
+
326
+ style = z.clone() #for adaptive groupnorm
327
+
328
+ noise = torch.randn_like(z).to(z.device) #generate noise
329
+ z = torch.cat([z, noise], dim=1) #concat noise to the style vector
330
+ z = self.conv_in(z)
331
+
332
+ ## mid
333
+ for res in range(self.num_res_blocks):
334
+ z = self.mid_block[res](z)
335
+
336
+ ## upsample
337
+ for i_level in reversed(range(self.num_blocks)):
338
+ ### pass in each resblock first adaGN
339
+ z = self.adaptive[i_level](z, style)
340
+ for i_block in range(self.num_res_blocks):
341
+ z = self.up[i_level].block[i_block](z)
342
+
343
+ if i_level > 0:
344
+ z = self.up[i_level].upsample(z)
345
+
346
+ z = self.norm_out(z)
347
+ z = swish(z)
348
+ z = self.conv_out(z)
349
+
350
+ return z
351
+
352
+
353
+ class VQModel(nn.Module):
354
+ def __init__(self,
355
+ ddconfig,
356
+ checkpoint=None,
357
+ gan_decoder = False,
358
+ ):
359
+ super().__init__()
360
+ self.encoder = Encoder(**ddconfig)
361
+ self.decoder = GANDecoder(**ddconfig) if gan_decoder else Decoder(**ddconfig)
362
+
363
+ # Load weights from the checkpoint
364
+ if checkpoint is not None:
365
+ self.load_from_ckpt(checkpoint)
366
+
367
+ def load_from_ckpt(self, checkpoint):
368
+ state = torch.load(checkpoint, mmap=True, map_location="cpu")
369
+ log_info = self.load_state_dict(state["state_dict"], strict=False)
370
+ has_missing_keys = bool(log_info.missing_keys)
371
+ has_unexpected_keys = bool(log_info.unexpected_keys)
372
+ if not has_missing_keys:
373
+ print(f"Successfully loaded all weights from checkpoint: {checkpoint}")
374
+ else:
375
+ if has_missing_keys:
376
+ print("Missing keys (model layers not in checkpoint):")
377
+ for key in log_info.missing_keys:
378
+ print(f" - {key}")
379
+ if False and has_unexpected_keys:
380
+ print("\nUnexpected keys (checkpoint layers not in model):")
381
+ for key in log_info.unexpected_keys:
382
+ print(f" - {key}")
383
+
384
+ def encode(self, x):
385
+ h = self.encoder(x)
386
+ codebook_value = torch.Tensor([1.0]).to(h)
387
+ quant_h = torch.where(h > 0, codebook_value, -codebook_value) # higher than 0 filled
388
+
389
+ return quant_h
390
+
391
+ # def vt_forward(self, image_list):
392
+ # q_list = []
393
+ # for x in image_list:
394
+ # quant = self.encode(x)
395
+ # quant = rearrange(quant.squeeze(0), "c h w -> (h w) c")
396
+ # q_list.append(quant)
397
+
398
+ # return torch.cat(q_list, dim=0)
399
+
400
+
401
+ def vt_forward(self, image_list, max_bs=32, ps=1):
402
+ groups = defaultdict(list) # {(H, W): [(idx, image_tensor), ...]}
403
+ for i, img in enumerate(image_list):
404
+ _, _, H, W = img.shape
405
+ groups[(H, W)].append((i, img))
406
+
407
+ output = [None] * len(image_list)
408
+
409
+ for (H, W), items in groups.items():
410
+ for start in range(0, len(items), max_bs):
411
+ chunk = items[start:start + max_bs]
412
+ idxs = [x[0] for x in chunk]
413
+ imgs = [x[1] for x in chunk]
414
+
415
+ batch = torch.cat(imgs, dim=0) # [B, 3, H, W]
416
+
417
+ quant = self.encode(batch) # [B, C, h, w]
418
+
419
+ for b in range(quant.size(0)):
420
+ q = rearrange(quant[b], "c (h p1) (w p2) -> (h w p1 p2) c", p1=ps, p2=ps)
421
+ output[idxs[b]] = q
422
+
423
+ return torch.cat(output, dim=0)
424
+
425
+ def vt_forward_maxpad(
426
+ self,
427
+ image_list,
428
+ max_bs=32,
429
+ stride=32,
430
+ min_size=256,
431
+ max_size=2048,
432
+ max_pixels=1024 * 1024,
433
+ normal_buckets=(384, 512, 768, 1024),
434
+ ):
435
+ """
436
+ image_list: list of [1, 3, H, W]
437
+ return: Tensor [(sum_i Hi*Wi/stride^2), C]
438
+ """
439
+
440
+ def is_long_image(H, W):
441
+ major = max(H, W)
442
+ minor = min(H, W)
443
+ return (
444
+ major >= 1024 and
445
+ minor <= 768 and
446
+ major / minor >= 1.5
447
+ )
448
+
449
+ groups = defaultdict(list)
450
+ sizes = {}
451
+
452
+ for idx, img in enumerate(image_list):
453
+ _, _, H, W = img.shape
454
+
455
+ # assert H >= min_size and W >= min_size
456
+ # assert H <= max_size and W <= max_size
457
+ # assert H * W <= max_pixels, f"image is too large: {H}x{W}"
458
+
459
+ if is_long_image(H, W):
460
+ bucket = "long"
461
+ else:
462
+ major = max(H, W)
463
+ for b in normal_buckets:
464
+ if major <= b:
465
+ bucket = b
466
+ break
467
+ else:
468
+ bucket = "long"
469
+
470
+ groups[bucket].append(idx)
471
+ sizes[idx] = (H, W)
472
+
473
+ output = [None] * len(image_list)
474
+
475
+
476
+ for bucket, idxs in groups.items():
477
+ imgs = [image_list[i] for i in idxs]
478
+
479
+ for start in range(0, len(imgs), max_bs):
480
+ batch_imgs = imgs[start:start + max_bs]
481
+ batch_idxs = idxs[start:start + max_bs]
482
+
483
+ H_max = max(img.shape[-2] for img in batch_imgs)
484
+ W_max = max(img.shape[-1] for img in batch_imgs)
485
+
486
+ H_pad = math.ceil(H_max / stride) * stride
487
+ W_pad = math.ceil(W_max / stride) * stride
488
+
489
+ padded = []
490
+ for img in batch_imgs:
491
+ _, _, H, W = img.shape
492
+ pad_h = H_pad - H
493
+ pad_w = W_pad - W
494
+ padded.append(F.pad(img, (0, pad_w, 0, pad_h)))
495
+
496
+ batch = torch.cat(padded, dim=0) # [B, 3, H_pad, W_pad]
497
+
498
+ quant = self.encode(batch) # [B, C, h', w']
499
+
500
+ for i, q in enumerate(quant):
501
+ H, W = sizes[batch_idxs[i]]
502
+ h_lat = math.ceil(H / stride)
503
+ w_lat = math.ceil(W / stride)
504
+
505
+ q = q[:, :h_lat, :w_lat]
506
+ q = rearrange(q, "c h w -> (h w) c")
507
+
508
+ output[batch_idxs[i]] = q
509
+
510
+ return torch.cat(output, dim=0)
511
+
512
+
513
+ def decode(self, quant):
514
+ dec = self.decoder(quant)
515
+ return dec
516
+
517
+ def forward(self, input):
518
+ quant = self.encode(input)
519
+ dec = self.decode(quant)
520
+ return dec, quant
modeling/vision_head/flow_head_parallel_x.py ADDED
@@ -0,0 +1,342 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import math
2
+
3
+ import torch
4
+ import torch.nn as nn
5
+ import torch.nn.functional as F
6
+
7
+ from .sampling_x import euler_maruyama
8
+
9
+ from flash_attn import flash_attn_func
10
+
11
+
12
+ def timestep_embedding(t, dim, max_period=10000, time_factor: float = 1000.0):
13
+ half = dim // 2
14
+ t = time_factor * t.float()
15
+ freqs = torch.exp(
16
+ -math.log(max_period)
17
+ * torch.arange(start=0, end=half, dtype=torch.float32, device=t.device)
18
+ / half
19
+ )
20
+
21
+ args = t[:, None] * freqs[None]
22
+ embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
23
+ if dim % 2:
24
+ embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1)
25
+ if torch.is_floating_point(t):
26
+ embedding = embedding.to(t)
27
+ return embedding
28
+
29
+ def time_shift_func(t: torch.Tensor, flow_shift: float = 1., sigma: float = 1.):
30
+ return (1 / flow_shift) / ( (1 / flow_shift) + (1 / t - 1) ** sigma)
31
+
32
+ class DiffHead(nn.Module):
33
+ def __init__(
34
+ self,
35
+ ch_target,
36
+ ch_cond,
37
+ ch_latent,
38
+ depth_latent,
39
+ depth_adanln,
40
+ grad_checkpointing=False,
41
+ time_shift=1.,
42
+ time_schedule='logit_normal',
43
+ P_mean=0.,
44
+ P_std=1.,
45
+ parallel_num=4,
46
+ diff_batch_mul=1,
47
+ use_swiglu=False,
48
+ ):
49
+ super(DiffHead, self).__init__()
50
+ self.ch_target = ch_target
51
+ self.time_shift = time_shift
52
+ self.time_schedule = time_schedule
53
+ self.P_mean = P_mean
54
+ self.P_std = P_std
55
+ self.diff_batch_mul = diff_batch_mul
56
+
57
+ self.net = TransEncoder(
58
+ in_channels=ch_target,
59
+ model_channels=ch_latent,
60
+ z_channels=ch_cond,
61
+ num_res_blocks=depth_latent,
62
+ num_ada_ln_blocks=depth_adanln,
63
+ grad_checkpointing=grad_checkpointing,
64
+ parallel_num=parallel_num,
65
+ use_swiglu=use_swiglu
66
+ )
67
+
68
+ def forward(self, x, cond):
69
+ with torch.autocast(device_type="cuda", enabled=False):
70
+ with torch.no_grad():
71
+ if self.time_schedule == 'logit_normal':
72
+ t = (torch.randn((x.shape[0]), device=x.device) * self.P_std + self.P_mean).sigmoid()
73
+ if self.time_shift != 1.:
74
+ t = time_shift_func(t, self.time_shift)
75
+ elif self.time_schedule == 'uniform':
76
+ t = torch.rand((x.shape[0]), device=x.device)
77
+ if self.time_shift != 1.:
78
+ t = time_shift_func(t, self.time_shift)
79
+ else:
80
+ raise NotImplementedError(f"unknown time_schedule {self.time_schedule}")
81
+ e = torch.randn_like(x)
82
+ ti = t.view(-1, 1, 1)
83
+ z = (1.0 - ti) * e + ti * x
84
+ v = (x - z) / (1 - ti).clamp_min(0.05)
85
+
86
+ if self.diff_batch_mul > 1:
87
+ chunks = self.diff_batch_mul
88
+ x_pred_list = []
89
+
90
+ z_chunks = torch.chunk(z, chunks, dim=0)
91
+ t_chunks = torch.chunk(t, chunks, dim=0)
92
+ cond_chunks = torch.chunk(cond, chunks, dim=0)
93
+ for z_i, t_i, cond_i in zip(z_chunks, t_chunks, cond_chunks):
94
+ output_i = self.net(z_i, t_i, cond_i)
95
+ x_pred_list.append(output_i)
96
+ x_pred = torch.cat(x_pred_list, dim=0)
97
+ else:
98
+ x_pred = self.net(z, t, cond)
99
+
100
+ v_pred = (x_pred - z) / (1 - ti).clamp_min(0.05)
101
+
102
+ with torch.autocast(device_type="cuda", enabled=False):
103
+ v_pred = v_pred.float()
104
+ loss = torch.mean((v - v_pred) ** 2, dim=2)
105
+ return loss
106
+
107
+ def sample(
108
+ self,
109
+ z,
110
+ cfg,
111
+ num_sampling_steps,
112
+ ):
113
+ return euler_maruyama(
114
+ self.ch_target,
115
+ self.net.forward,
116
+ z,
117
+ cfg,
118
+ num_sampling_steps=num_sampling_steps,
119
+ time_shift = self.time_shift,
120
+ )
121
+
122
+ def initialize_weights(self):
123
+ self.net.initialize_weights()
124
+
125
+
126
+ class TimestepEmbedder(nn.Module):
127
+ """
128
+ Embeds scalar timesteps into vector representations.
129
+ """
130
+
131
+ def __init__(self, hidden_size, frequency_embedding_size=256):
132
+ super().__init__()
133
+ self.mlp = nn.Sequential(
134
+ nn.Linear(frequency_embedding_size, hidden_size, bias=True),
135
+ nn.SiLU(),
136
+ nn.Linear(hidden_size, hidden_size, bias=True),
137
+ )
138
+ self.frequency_embedding_size = frequency_embedding_size
139
+
140
+ def forward(self, t):
141
+ t_freq = timestep_embedding(t, self.frequency_embedding_size)
142
+ t_emb = self.mlp(t_freq)
143
+ return t_emb
144
+
145
+
146
+ class ResBlock(nn.Module):
147
+ def __init__(self, channels):
148
+ super().__init__()
149
+ self.channels = channels
150
+ self.norm = nn.LayerNorm(channels, eps=1e-6, elementwise_affine=True)
151
+ hidden_dim = int(channels * 1.5)
152
+ self.w1 = nn.Linear(channels, hidden_dim * 2, bias=True)
153
+ self.w2 = nn.Linear(hidden_dim, channels, bias=True)
154
+
155
+ def forward(self, x, scale, shift, gate):
156
+ h = self.norm(x) * (1 + scale) + shift
157
+ h1, h2 = self.w1(h).chunk(2, dim=-1)
158
+ h = self.w2(F.silu(h1) * h2)
159
+ return x + h * gate
160
+
161
+
162
+ class FinalLayer(nn.Module):
163
+ def __init__(self, channels, out_channels):
164
+ super().__init__()
165
+ self.norm_final = nn.LayerNorm(channels, eps=1e-6, elementwise_affine=False)
166
+ self.ada_ln_modulation = nn.Linear(channels, channels * 2, bias=True)
167
+ self.linear = nn.Linear(channels, out_channels, bias=True)
168
+
169
+ def forward(self, x, y):
170
+ scale, shift = self.ada_ln_modulation(y).chunk(2, dim=-1)
171
+ x = self.norm_final(x) * (1.0 + scale) + shift
172
+ x = self.linear(x)
173
+ return x
174
+
175
+ class Attention(nn.Module):
176
+ def __init__(
177
+ self,
178
+ dim,
179
+ n_head,
180
+ ):
181
+ super().__init__()
182
+ assert dim % n_head == 0
183
+ self.dim = dim
184
+ self.head_dim = dim // n_head
185
+ self.scale = self.head_dim**-0.5
186
+ self.n_head = n_head
187
+ total_kv_dim = (self.n_head * 3) * self.head_dim
188
+
189
+ self.wqkv = nn.Linear(dim, total_kv_dim, bias=True)
190
+ self.wo = nn.Linear(dim, dim, bias=True)
191
+
192
+ def forward(
193
+ self,
194
+ x: torch.Tensor,
195
+ ):
196
+ bsz, seqlen, _ = x.shape
197
+ xq, xk, xv = self.wqkv(x).chunk(3, dim=-1)
198
+
199
+ xq = xq.view(bsz, seqlen, self.n_head, self.head_dim)
200
+ xk = xk.view(bsz, seqlen, self.n_head, self.head_dim)
201
+ xv = xv.view(bsz, seqlen, self.n_head, self.head_dim)
202
+
203
+ if seqlen <= 32:
204
+ xq, xk, xv = map(lambda x: x.transpose(1, 2), (xq, xk, xv))
205
+ xq = xq * self.scale
206
+ attn = xq @ xk.transpose(-1, -2)
207
+ attn = F.softmax(attn, dim=-1)
208
+ output = (attn @ xv).transpose(1, 2).contiguous()
209
+ else:
210
+ output = flash_attn_func(
211
+ xq,
212
+ xk,
213
+ xv,
214
+ causal=False,
215
+ )
216
+
217
+ output = output.view(bsz, seqlen, self.dim)
218
+
219
+ output = self.wo(output)
220
+ return output
221
+
222
+ class TransBlock(nn.Module):
223
+ def __init__(self, channels, use_swiglu=False):
224
+ super().__init__()
225
+ self.channels = channels
226
+ self.norm1 = nn.LayerNorm(channels, eps=1e-6, elementwise_affine=True)
227
+ self.attn = Attention(channels, n_head=channels // 128)
228
+
229
+ self.norm2 = nn.LayerNorm(channels, eps=1e-6, elementwise_affine=True)
230
+ hidden_dim = int(channels * 1.5)
231
+ self.use_swiglu = use_swiglu
232
+ if not self.use_swiglu:
233
+ self.mlp = nn.Sequential(
234
+ nn.Linear(self.channels, hidden_dim),
235
+ nn.SiLU(),
236
+ nn.Linear(hidden_dim, self.channels),
237
+ )
238
+ else:
239
+ self.w1 = nn.Linear(channels, hidden_dim * 2, bias=True)
240
+ self.w2 = nn.Linear(hidden_dim, channels, bias=True)
241
+
242
+ def forward(self, x, scale1, shift1, gate1, scale2, shift2, gate2):
243
+ h = self.norm1(x) * (1 + scale1) + shift1
244
+ h = self.attn(h)
245
+ x = x + h * gate1
246
+ h = self.norm2(x) * (1 + scale2) + shift2
247
+ if not self.use_swiglu:
248
+ h = self.mlp(h)
249
+ else:
250
+ h1, h2 = self.w1(h).chunk(2, dim=-1)
251
+ h = self.w2(F.silu(h1) * h2)
252
+ return x + h * gate2
253
+
254
+ class TransEncoder(nn.Module):
255
+
256
+ def __init__(
257
+ self,
258
+ in_channels,
259
+ model_channels,
260
+ z_channels,
261
+ num_res_blocks,
262
+ num_ada_ln_blocks=2,
263
+ grad_checkpointing=False,
264
+ parallel_num=4,
265
+ use_swiglu=False,
266
+ ):
267
+ super().__init__()
268
+
269
+ self.in_channels = in_channels
270
+ self.model_channels = model_channels
271
+ self.out_channels = in_channels
272
+ self.num_res_blocks = num_res_blocks
273
+ self.grad_checkpointing = grad_checkpointing
274
+ self.parallel_num = parallel_num
275
+
276
+ self.time_embed = TimestepEmbedder(model_channels)
277
+ self.cond_embed = nn.Linear(z_channels, model_channels)
278
+
279
+ self.input_proj = nn.Linear(in_channels, model_channels)
280
+ self.res_blocks = nn.ModuleList()
281
+ for i in range(num_res_blocks):
282
+ self.res_blocks.append(
283
+ TransBlock(
284
+ model_channels,
285
+ use_swiglu
286
+ )
287
+ )
288
+ # share adaLN for consecutive blocks, to save computation and parameters
289
+ self.ada_ln_blocks = nn.ModuleList()
290
+ for i in range(num_ada_ln_blocks):
291
+ self.ada_ln_blocks.append(
292
+ nn.Linear(model_channels, model_channels * 6, bias=True)
293
+ )
294
+ self.ada_ln_switch_freq = max(1, num_res_blocks // num_ada_ln_blocks)
295
+ assert (
296
+ num_res_blocks % self.ada_ln_switch_freq
297
+ ) == 0, "num_res_blocks must be divisible by num_ada_ln_blocks"
298
+ self.final_layer = FinalLayer(model_channels, self.out_channels)
299
+
300
+ self.initialize_weights()
301
+
302
+ def initialize_weights(self):
303
+ def _basic_init(module):
304
+ if isinstance(module, nn.Linear):
305
+ torch.nn.init.xavier_uniform_(module.weight)
306
+ if module.bias is not None:
307
+ nn.init.constant_(module.bias, 0)
308
+
309
+ self.apply(_basic_init)
310
+
311
+ # Initialize timestep embedding MLP
312
+ nn.init.normal_(self.time_embed.mlp[0].weight, std=0.02)
313
+ nn.init.normal_(self.time_embed.mlp[2].weight, std=0.02)
314
+
315
+ for block in self.ada_ln_blocks:
316
+ nn.init.constant_(block.weight, 0)
317
+ nn.init.constant_(block.bias, 0)
318
+
319
+ # Zero-out output layers
320
+ nn.init.constant_(self.final_layer.ada_ln_modulation.weight, 0)
321
+ nn.init.constant_(self.final_layer.ada_ln_modulation.bias, 0)
322
+ nn.init.constant_(self.final_layer.linear.weight, 0)
323
+ nn.init.constant_(self.final_layer.linear.bias, 0)
324
+
325
+ def forward(self, x, t, c):
326
+ x = self.input_proj(x)
327
+ t = self.time_embed(t).unsqueeze(1)
328
+ c = self.cond_embed(c)
329
+
330
+ y = F.silu(t + c)
331
+ scale1, shift1, gate1, scale2, shift2, gate2 = self.ada_ln_blocks[0](y).chunk(6, dim=-1)
332
+
333
+ for i, block in enumerate(self.res_blocks):
334
+ if i > 0 and i % self.ada_ln_switch_freq == 0:
335
+ ada_ln_block = self.ada_ln_blocks[i // self.ada_ln_switch_freq]
336
+ scale1, shift1, gate1, scale2, shift2, gate2 = ada_ln_block(y).chunk(6, dim=-1)
337
+ x = block(x, scale1, shift1, gate1, scale2, shift2, gate2)
338
+
339
+ output = self.final_layer(x, y)
340
+
341
+ # use sigmoid to map to [-1, 1]
342
+ return 2 * torch.sigmoid(output) - 1
modeling/vision_head/sampling_x.py ADDED
@@ -0,0 +1,125 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+
3
+ def time_shift_func(t: torch.Tensor, flow_shift: float = 1., sigma: float = 1.):
4
+ return (1 / flow_shift) / ( (1 / flow_shift) + (1 / t - 1) ** sigma)
5
+
6
+ def get_score_from_velocity(velocity, x, t):
7
+ alpha_t, d_alpha_t = t, 1
8
+ sigma_t, d_sigma_t = 1 - t, -1
9
+ mean = x
10
+ reverse_alpha_ratio = alpha_t / d_alpha_t
11
+ var = sigma_t**2 - reverse_alpha_ratio * d_sigma_t * sigma_t
12
+ score = (reverse_alpha_ratio * velocity - mean) / var
13
+ return score
14
+
15
+
16
+ def get_velocity_from_cfg(velocity, cfg, cfg_mult):
17
+ if cfg_mult == 2:
18
+ cond_v, uncond_v = torch.chunk(velocity, 2, dim=0)
19
+ velocity = uncond_v + cfg * (cond_v - uncond_v)
20
+ return velocity
21
+
22
+
23
+ # @torch.compile()
24
+ def euler_step(x, v, dt: float, cfg: float, cfg_mult: int):
25
+ with torch.amp.autocast("cuda", enabled=False):
26
+ v = v.to(torch.float32)
27
+ v = get_velocity_from_cfg(v, cfg, cfg_mult)
28
+ x = x + v * dt
29
+ return x
30
+
31
+
32
+ # @torch.compile()
33
+ def euler_maruyama_step(x, v, t, dt: float, cfg: float, cfg_mult: int):
34
+ with torch.amp.autocast("cuda", enabled=False):
35
+ v = v.to(torch.float32)
36
+ v = get_velocity_from_cfg(v, cfg, cfg_mult)
37
+ score = get_score_from_velocity(v, x, t)
38
+ drift = v + (1 - t) * score
39
+ noise_scale = (2.0 * (1.0 - t) * dt) ** 0.5
40
+ x = x + drift * dt + noise_scale * torch.randn_like(x)
41
+ return x
42
+
43
+
44
+ def euler_maruyama(
45
+ input_dim,
46
+ forward_fn,
47
+ c: torch.Tensor,
48
+ cfg: float = 1.0,
49
+ num_sampling_steps: int = 20,
50
+ last_step_size: float = 0.05,
51
+ time_shift: float = 1.,
52
+ ):
53
+ cfg_mult = 1
54
+ if cfg > 1.0:
55
+ cfg_mult += 1
56
+
57
+ x_shape = list(c.shape)
58
+ x_shape[0] = x_shape[0] // cfg_mult
59
+ x_shape[-1] = input_dim
60
+ x = torch.randn(x_shape, device=c.device)
61
+ # an = (1.0 - last_step_size) / num_sampling_steps
62
+ t_all = torch.linspace(0, 1-last_step_size, num_sampling_steps+1, device=c.device, dtype=torch.float32)
63
+ t_all = time_shift_func(t_all, time_shift)
64
+ dt = t_all[1:] - t_all[:-1]
65
+ t = torch.tensor(
66
+ 0.0, device=c.device, dtype=torch.float32
67
+ ) # use tensor to avoid compile warning
68
+ t_batch = torch.zeros(c.shape[0], device=c.device)
69
+ for i in range(num_sampling_steps):
70
+ t_batch[:] = t
71
+ combined = torch.cat([x] * cfg_mult, dim=0)
72
+ output = forward_fn(
73
+ combined,
74
+ t_batch,
75
+ c,
76
+ )
77
+ if output.dim() == 2:
78
+ v = (output - combined) / (1 - t_batch.view(-1,1)).clamp_min(0.05)
79
+ elif output.dim() == 3:
80
+ v = (output - combined) / (1 - t_batch.view(-1,1,1)).clamp_min(0.05)
81
+ x = euler_maruyama_step(x, v, t, dt[i], cfg, cfg_mult)
82
+ t += dt[i]
83
+
84
+ combined = torch.cat([x] * cfg_mult, dim=0)
85
+ t_batch[:] = 1 - last_step_size
86
+ output = forward_fn(
87
+ combined,
88
+ t_batch,
89
+ c,
90
+ )
91
+ if output.dim() == 2:
92
+ v = (output - combined) / (1 - t_batch.view(-1,1)).clamp_min(0.05)
93
+ elif output.dim() == 3:
94
+ v = (output - combined) / (1 - t_batch.view(-1,1,1)).clamp_min(0.05)
95
+ x = euler_step(x, v, last_step_size, cfg, cfg_mult)
96
+
97
+ return torch.cat([x] * cfg_mult, dim=0)
98
+
99
+
100
+ def euler(
101
+ input_dim,
102
+ forward_fn,
103
+ c,
104
+ cfg: float = 1.0,
105
+ num_sampling_steps: int = 50,
106
+ ):
107
+ cfg_mult = 1
108
+ if cfg > 1.0:
109
+ cfg_mult = 2
110
+
111
+ x_shape = list(c.shape)
112
+ x_shape[0] = x_shape[0] // cfg_mult
113
+ x_shape[-1] = input_dim
114
+ x = torch.randn(x_shape, device=c.device)
115
+ dt = 1.0 / num_sampling_steps
116
+ t = 0
117
+ t_batch = torch.zeros(c.shape[0], device=c.device)
118
+ for _ in range(num_sampling_steps):
119
+ t_batch[:] = t
120
+ combined = torch.cat([x] * cfg_mult, dim=0)
121
+ v = forward_fn(combined, t_batch, c)
122
+ x = euler_step(x, v, dt, cfg, cfg_mult)
123
+ t += dt
124
+
125
+ return torch.cat([x] * cfg_mult, dim=0)
requirements.txt ADDED
@@ -0,0 +1,10 @@
 
 
 
 
 
 
 
 
 
 
 
1
+ transformers==4.57.0
2
+ omegaconf
3
+ liger-kernel
4
+ numpy==1.26.4
5
+ huggingface_hub==0.34.4
6
+ einops==0.6.1
7
+ torch==2.7.1
8
+ torchvision==0.22.1
9
+ safetensors==0.6.2
10
+ gradio