Spaces:
Paused
Paused
init space
Browse files- .gitignore +16 -0
- .gradio/certificate.pem +31 -0
- app.py +231 -0
- modeling/t2i_pipeline.py +283 -0
- modeling/utils.py +216 -0
- modeling/vision_encoder/autoencoder.py +520 -0
- modeling/vision_head/flow_head_parallel_x.py +342 -0
- modeling/vision_head/sampling_x.py +125 -0
- requirements.txt +10 -0
.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 |
+
[](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
|