luca115 commited on
Commit
3a75c87
·
verified ·
1 Parent(s): 1416514

Upload folder using huggingface_hub

Browse files
This view is limited to 50 files because it contains too many changes.   See raw diff
Files changed (50) hide show
  1. LICENSE +201 -0
  2. README.md +15 -7
  3. app.py +258 -0
  4. common/__init__.py +0 -0
  5. common/cache.py +47 -0
  6. common/config.py +110 -0
  7. common/decorators.py +147 -0
  8. common/diffusion/__init__.py +56 -0
  9. common/diffusion/config.py +74 -0
  10. common/diffusion/samplers/base.py +108 -0
  11. common/diffusion/samplers/euler.py +89 -0
  12. common/diffusion/schedules/base.py +131 -0
  13. common/diffusion/schedules/lerp.py +55 -0
  14. common/diffusion/timesteps/base.py +72 -0
  15. common/diffusion/timesteps/sampling/trailing.py +49 -0
  16. common/diffusion/types.py +59 -0
  17. common/diffusion/utils.py +84 -0
  18. common/distributed/__init__.py +37 -0
  19. common/distributed/advanced.py +208 -0
  20. common/distributed/basic.py +84 -0
  21. common/distributed/meta_init_utils.py +41 -0
  22. common/distributed/ops.py +494 -0
  23. common/logger.py +44 -0
  24. common/partition.py +59 -0
  25. common/seed.py +30 -0
  26. configs_3b/main.yaml +88 -0
  27. data/image/transforms/area_resize.py +135 -0
  28. data/image/transforms/divisible_crop.py +40 -0
  29. data/image/transforms/na_resize.py +50 -0
  30. data/image/transforms/side_resize.py +54 -0
  31. data/video/transforms/rearrange.py +24 -0
  32. models/dit_v2/attention.py +86 -0
  33. models/dit_v2/embedding.py +62 -0
  34. models/dit_v2/mlp.py +62 -0
  35. models/dit_v2/mm.py +74 -0
  36. models/dit_v2/modulation.py +102 -0
  37. models/dit_v2/na.py +241 -0
  38. models/dit_v2/nablocks/__init__.py +26 -0
  39. models/dit_v2/nablocks/attention/__init__.py +25 -0
  40. models/dit_v2/nablocks/attention/mmattn.py +266 -0
  41. models/dit_v2/nablocks/mmsr_block.py +119 -0
  42. models/dit_v2/nadit.py +246 -0
  43. models/dit_v2/normalization.py +63 -0
  44. models/dit_v2/patch/__init__.py +19 -0
  45. models/dit_v2/patch/patch_v1.py +127 -0
  46. models/dit_v2/rope.py +150 -0
  47. models/dit_v2/window.py +83 -0
  48. models/video_vae_v3/modules/attn_video_vae.py +1345 -0
  49. models/video_vae_v3/modules/causal_inflation_lib.py +460 -0
  50. models/video_vae_v3/modules/context_parallel_lib.py +164 -0
LICENSE ADDED
@@ -0,0 +1,201 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ Apache License
2
+ Version 2.0, January 2004
3
+ http://www.apache.org/licenses/
4
+
5
+ TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
6
+
7
+ 1. Definitions.
8
+
9
+ "License" shall mean the terms and conditions for use, reproduction,
10
+ and distribution as defined by Sections 1 through 9 of this document.
11
+
12
+ "Licensor" shall mean the copyright owner or entity authorized by
13
+ the copyright owner that is granting the License.
14
+
15
+ "Legal Entity" shall mean the union of the acting entity and all
16
+ other entities that control, are controlled by, or are under common
17
+ control with that entity. For the purposes of this definition,
18
+ "control" means (i) the power, direct or indirect, to cause the
19
+ direction or management of such entity, whether by contract or
20
+ otherwise, or (ii) ownership of fifty percent (50%) or more of the
21
+ outstanding shares, or (iii) beneficial ownership of such entity.
22
+
23
+ "You" (or "Your") shall mean an individual or Legal Entity
24
+ exercising permissions granted by this License.
25
+
26
+ "Source" form shall mean the preferred form for making modifications,
27
+ including but not limited to software source code, documentation
28
+ source, and configuration files.
29
+
30
+ "Object" form shall mean any form resulting from mechanical
31
+ transformation or translation of a Source form, including but
32
+ not limited to compiled object code, generated documentation,
33
+ and conversions to other media types.
34
+
35
+ "Work" shall mean the work of authorship, whether in Source or
36
+ Object form, made available under the License, as indicated by a
37
+ copyright notice that is included in or attached to the work
38
+ (an example is provided in the Appendix below).
39
+
40
+ "Derivative Works" shall mean any work, whether in Source or Object
41
+ form, that is based on (or derived from) the Work and for which the
42
+ editorial revisions, annotations, elaborations, or other modifications
43
+ represent, as a whole, an original work of authorship. For the purposes
44
+ of this License, Derivative Works shall not include works that remain
45
+ separable from, or merely link (or bind by name) to the interfaces of,
46
+ the Work and Derivative Works thereof.
47
+
48
+ "Contribution" shall mean any work of authorship, including
49
+ the original version of the Work and any modifications or additions
50
+ to that Work or Derivative Works thereof, that is intentionally
51
+ submitted to Licensor for inclusion in the Work by the copyright owner
52
+ or by an individual or Legal Entity authorized to submit on behalf of
53
+ the copyright owner. For the purposes of this definition, "submitted"
54
+ means any form of electronic, verbal, or written communication sent
55
+ to the Licensor or its representatives, including but not limited to
56
+ communication on electronic mailing lists, source code control systems,
57
+ and issue tracking systems that are managed by, or on behalf of, the
58
+ Licensor for the purpose of discussing and improving the Work, but
59
+ excluding communication that is conspicuously marked or otherwise
60
+ designated in writing by the copyright owner as "Not a Contribution."
61
+
62
+ "Contributor" shall mean Licensor and any individual or Legal Entity
63
+ on behalf of whom a Contribution has been received by Licensor and
64
+ subsequently incorporated within the Work.
65
+
66
+ 2. Grant of Copyright License. Subject to the terms and conditions of
67
+ this License, each Contributor hereby grants to You a perpetual,
68
+ worldwide, non-exclusive, no-charge, royalty-free, irrevocable
69
+ copyright license to reproduce, prepare Derivative Works of,
70
+ publicly display, publicly perform, sublicense, and distribute the
71
+ Work and such Derivative Works in Source or Object form.
72
+
73
+ 3. Grant of Patent License. Subject to the terms and conditions of
74
+ this License, each Contributor hereby grants to You a perpetual,
75
+ worldwide, non-exclusive, no-charge, royalty-free, irrevocable
76
+ (except as stated in this section) patent license to make, have made,
77
+ use, offer to sell, sell, import, and otherwise transfer the Work,
78
+ where such license applies only to those patent claims licensable
79
+ by such Contributor that are necessarily infringed by their
80
+ Contribution(s) alone or by combination of their Contribution(s)
81
+ with the Work to which such Contribution(s) was submitted. If You
82
+ institute patent litigation against any entity (including a
83
+ cross-claim or counterclaim in a lawsuit) alleging that the Work
84
+ or a Contribution incorporated within the Work constitutes direct
85
+ or contributory patent infringement, then any patent licenses
86
+ granted to You under this License for that Work shall terminate
87
+ as of the date such litigation is filed.
88
+
89
+ 4. Redistribution. You may reproduce and distribute copies of the
90
+ Work or Derivative Works thereof in any medium, with or without
91
+ modifications, and in Source or Object form, provided that You
92
+ meet the following conditions:
93
+
94
+ (a) You must give any other recipients of the Work or
95
+ Derivative Works a copy of this License; and
96
+
97
+ (b) You must cause any modified files to carry prominent notices
98
+ stating that You changed the files; and
99
+
100
+ (c) You must retain, in the Source form of any Derivative Works
101
+ that You distribute, all copyright, patent, trademark, and
102
+ attribution notices from the Source form of the Work,
103
+ excluding those notices that do not pertain to any part of
104
+ the Derivative Works; and
105
+
106
+ (d) If the Work includes a "NOTICE" text file as part of its
107
+ distribution, then any Derivative Works that You distribute must
108
+ include a readable copy of the attribution notices contained
109
+ within such NOTICE file, excluding those notices that do not
110
+ pertain to any part of the Derivative Works, in at least one
111
+ of the following places: within a NOTICE text file distributed
112
+ as part of the Derivative Works; within the Source form or
113
+ documentation, if provided along with the Derivative Works; or,
114
+ within a display generated by the Derivative Works, if and
115
+ wherever such third-party notices normally appear. The contents
116
+ of the NOTICE file are for informational purposes only and
117
+ do not modify the License. You may add Your own attribution
118
+ notices within Derivative Works that You distribute, alongside
119
+ or as an addendum to the NOTICE text from the Work, provided
120
+ that such additional attribution notices cannot be construed
121
+ as modifying the License.
122
+
123
+ You may add Your own copyright statement to Your modifications and
124
+ may provide additional or different license terms and conditions
125
+ for use, reproduction, or distribution of Your modifications, or
126
+ for any such Derivative Works as a whole, provided Your use,
127
+ reproduction, and distribution of the Work otherwise complies with
128
+ the conditions stated in this License.
129
+
130
+ 5. Submission of Contributions. Unless You explicitly state otherwise,
131
+ any Contribution intentionally submitted for inclusion in the Work
132
+ by You to the Licensor shall be under the terms and conditions of
133
+ this License, without any additional terms or conditions.
134
+ Notwithstanding the above, nothing herein shall supersede or modify
135
+ the terms of any separate license agreement you may have executed
136
+ with Licensor regarding such Contributions.
137
+
138
+ 6. Trademarks. This License does not grant permission to use the trade
139
+ names, trademarks, service marks, or product names of the Licensor,
140
+ except as required for reasonable and customary use in describing the
141
+ origin of the Work and reproducing the content of the NOTICE file.
142
+
143
+ 7. Disclaimer of Warranty. Unless required by applicable law or
144
+ agreed to in writing, Licensor provides the Work (and each
145
+ Contributor provides its Contributions) on an "AS IS" BASIS,
146
+ WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
147
+ implied, including, without limitation, any warranties or conditions
148
+ of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
149
+ PARTICULAR PURPOSE. You are solely responsible for determining the
150
+ appropriateness of using or redistributing the Work and assume any
151
+ risks associated with Your exercise of permissions under this License.
152
+
153
+ 8. Limitation of Liability. In no event and under no legal theory,
154
+ whether in tort (including negligence), contract, or otherwise,
155
+ unless required by applicable law (such as deliberate and grossly
156
+ negligent acts) or agreed to in writing, shall any Contributor be
157
+ liable to You for damages, including any direct, indirect, special,
158
+ incidental, or consequential damages of any character arising as a
159
+ result of this License or out of the use or inability to use the
160
+ Work (including but not limited to damages for loss of goodwill,
161
+ work stoppage, computer failure or malfunction, or any and all
162
+ other commercial damages or losses), even if such Contributor
163
+ has been advised of the possibility of such damages.
164
+
165
+ 9. Accepting Warranty or Additional Liability. While redistributing
166
+ the Work or Derivative Works thereof, You may choose to offer,
167
+ and charge a fee for, acceptance of support, warranty, indemnity,
168
+ or other liability obligations and/or rights consistent with this
169
+ License. However, in accepting such obligations, You may act only
170
+ on Your own behalf and on Your sole responsibility, not on behalf
171
+ of any other Contributor, and only if You agree to indemnify,
172
+ defend, and hold each Contributor harmless for any liability
173
+ incurred by, or claims asserted against, such Contributor by reason
174
+ of your accepting any such warranty or additional liability.
175
+
176
+ END OF TERMS AND CONDITIONS
177
+
178
+ APPENDIX: How to apply the Apache License to your work.
179
+
180
+ To apply the Apache License to your work, attach the following
181
+ boilerplate notice, with the fields enclosed by brackets "{}"
182
+ replaced with your own identifying information. (Don't include
183
+ the brackets!) The text should be enclosed in the appropriate
184
+ comment syntax for the file format. We also recommend that a
185
+ file or class name and description of purpose be included on the
186
+ same "printed page" as the copyright notice for easier
187
+ identification within third-party archives.
188
+
189
+ Copyright 2025 seed
190
+
191
+ Licensed under the Apache License, Version 2.0 (the "License");
192
+ you may not use this file except in compliance with the License.
193
+ You may obtain a copy of the License at
194
+
195
+ http://www.apache.org/licenses/LICENSE-2.0
196
+
197
+ Unless required by applicable law or agreed to in writing, software
198
+ distributed under the License is distributed on an "AS IS" BASIS,
199
+ WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
200
+ See the License for the specific language governing permissions and
201
+ limitations under the License.
README.md CHANGED
@@ -1,13 +1,21 @@
1
  ---
2
- title: Seedvr2 3b
3
- emoji: 🏢
4
- colorFrom: yellow
5
- colorTo: gray
6
  sdk: gradio
7
- sdk_version: 6.24.0
8
- python_version: '3.13'
9
  app_file: app.py
10
  pinned: false
 
 
 
 
11
  ---
12
 
13
- Check out the configuration reference at https://huggingface.co/docs/hub/spaces-config-reference
 
 
 
 
 
 
1
  ---
2
+ title: SeedVR2 3B Image Upscaler
3
+ emoji: 🔍
4
+ colorFrom: blue
5
+ colorTo: purple
6
  sdk: gradio
7
+ sdk_version: 5.49.1
 
8
  app_file: app.py
9
  pinned: false
10
+ license: apache-2.0
11
+ short_description: One-step diffusion image restoration with SeedVR2-3B
12
+ models:
13
+ - ByteDance-Seed/SeedVR2-3B
14
  ---
15
 
16
+ # SeedVR2-3B Image Upscaler
17
+
18
+ One-step diffusion restoration for images, from
19
+ [ByteDance-Seed/SeedVR2-3B](https://huggingface.co/ByteDance-Seed/SeedVR2-3B)
20
+ (Apache-2.0). This Space serves the image path only; see the header of `app.py`
21
+ for what was changed from the reference implementation and why.
app.py ADDED
@@ -0,0 +1,258 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SeedVR2-3B image restoration, Upsampler v3 recipe.
2
+ #
3
+ # Ported from ByteDance-Seed/SeedVR2-3B. What changed and why:
4
+ #
5
+ # 1. Image only. The upstream Space serves video and images from one entry
6
+ # point, carrying the sequence-parallel plumbing, frame cutting and video
7
+ # writing along with it. This Space backs an image tool, so that is all gone.
8
+ # 2. ONE @spaces.GPU entry. Upstream decorates configure_runner,
9
+ # generation_step AND generation_loop, so a single request booked three GPU
10
+ # allocations of 100s each. ZeroGPU checks the requested duration against the
11
+ # visitor's remaining quota, and an unauthenticated visitor has 120 seconds a
12
+ # day in total, so the upstream shape cannot serve an anonymous user at all.
13
+ # Everything now runs inside one call with a measured dynamic duration.
14
+ # 3. No apex. Upstream installs a prebuilt `apex-0.1-cp310-...whl` and selects
15
+ # `fusedrms` / `fusedln` norms in configs_3b/main.yaml. ZeroGPU runs Python
16
+ # 3.12, where that wheel does not install, so every norm layer then failed.
17
+ # The config now selects the `rms` / `layer` paths that the same source file
18
+ # already implements in pure PyTorch, with identical parameter names and
19
+ # shapes so the checkpoint loads unchanged.
20
+ # 4. No hard flash-attn dependency. See models/dit_v2/attention.py.
21
+ #
22
+ # The model restores at a fixed ~3.7MP working resolution regardless of input
23
+ # size (it was trained at high res and NaResize scales the input to meet it), so
24
+ # GPU cost per request is essentially constant and no tiling is involved.
25
+
26
+ import gc
27
+ import os
28
+ import mimetypes
29
+ from pathlib import Path
30
+
31
+ import gradio as gr
32
+ import spaces
33
+ import torch
34
+ import torch.nn.functional as F
35
+ from einops import rearrange
36
+ from omegaconf import OmegaConf
37
+ from PIL import Image
38
+ from huggingface_hub import hf_hub_download
39
+ from torchvision.transforms import Compose, Lambda, Normalize
40
+ import torchvision.transforms as T
41
+
42
+ from upsampler_theme import UPSAMPLER_CSS, UPSAMPLER_THEME, footer_html, header_html
43
+
44
+ from data.image.transforms.divisible_crop import DivisibleCrop
45
+ from data.image.transforms.na_resize import NaResize
46
+ from data.video.transforms.rearrange import Rearrange
47
+ from common.config import load_config
48
+ from common.distributed import init_torch
49
+ from common.seed import set_seed
50
+ from projects.video_diffusion_sr.infer import VideoDiffusionInfer
51
+
52
+ try:
53
+ from projects.video_diffusion_sr.color_fix import wavelet_reconstruction
54
+
55
+ USE_COLOR_FIX = True
56
+ except ImportError:
57
+ USE_COLOR_FIX = False
58
+ print("color fix unavailable; output will not be wavelet-reconstructed")
59
+
60
+ # Weights come from the official ByteDance repo rather than a re-upload. It is
61
+ # the canonical source for these files and is Apache-2.0, so there is no mirror
62
+ # in the chain that could change under us.
63
+ WEIGHTS_REPO = "ByteDance-Seed/SeedVR2-3B"
64
+ CKPT_DIR = Path("./ckpts")
65
+ CKPT_DIR.mkdir(exist_ok=True)
66
+
67
+
68
+ def _fetch(filename: str, target: Path) -> str:
69
+ if target.exists():
70
+ return str(target)
71
+ path = hf_hub_download(repo_id=WEIGHTS_REPO, filename=filename)
72
+ target.symlink_to(path)
73
+ return str(target)
74
+
75
+
76
+ DIT_CKPT = _fetch("seedvr2_ema_3b.pth", CKPT_DIR / "seedvr2_ema_3b.pth")
77
+ VAE_CKPT = _fetch("ema_vae.pth", CKPT_DIR / "ema_vae.pth")
78
+ # The text branch is conditioned by two fixed embeddings shipped with the
79
+ # weights, so this Space runs no text encoder at all.
80
+ POS_EMB = _fetch("pos_emb.pt", Path("./pos_emb.pt"))
81
+ NEG_EMB = _fetch("neg_emb.pt", Path("./neg_emb.pt"))
82
+
83
+ # The resolution the model restores at, as an area. Upstream hardcodes
84
+ # 2560*1440 for images with `downsample_only=False`, meaning small inputs are
85
+ # scaled UP to it and large inputs down, because the model was only trained at
86
+ # high resolution. Keeping that exact value keeps output quality identical to
87
+ # the reference implementation.
88
+ WORK_AREA = 2560 * 1440
89
+
90
+ # Single-process "distributed" context. The model code routes every device
91
+ # placement through common.distributed, which reads these.
92
+ os.environ.setdefault("MASTER_ADDR", "127.0.0.1")
93
+ os.environ.setdefault("MASTER_PORT", "12355")
94
+ os.environ.setdefault("RANK", "0")
95
+ os.environ.setdefault("WORLD_SIZE", "1")
96
+ os.environ.setdefault("LOCAL_RANK", "0")
97
+
98
+ _runner = None
99
+
100
+
101
+ def _ensure_runner():
102
+ """Build the runner once, inside a GPU context.
103
+
104
+ Deliberately NOT done at module scope, even though ZeroGPU prefers that for
105
+ placement: `init_torch` ends in `dist.init_process_group(backend="nccl")`
106
+ and `torch.cuda.set_device`, which need a real device rather than the CUDA
107
+ emulation that applies outside `@spaces.GPU`. Memoized because
108
+ `init_process_group` raises if called twice, and because a warm worker
109
+ should not reload 3B parameters per request.
110
+ """
111
+ global _runner
112
+ if _runner is not None:
113
+ return _runner
114
+
115
+ if not torch.distributed.is_initialized():
116
+ init_torch(cudnn_benchmark=False)
117
+
118
+ runner = VideoDiffusionInfer(load_config(os.path.join("./configs_3b", "main.yaml")))
119
+ OmegaConf.set_readonly(runner.config, False)
120
+ runner.configure_dit_model(device="cuda", checkpoint=DIT_CKPT)
121
+ runner.configure_vae_model()
122
+ if hasattr(runner.vae, "set_memory_limit"):
123
+ runner.vae.set_memory_limit(**runner.config.vae.memory_limit)
124
+
125
+ _runner = runner
126
+ return _runner
127
+
128
+
129
+ def _transform():
130
+ return Compose(
131
+ [
132
+ NaResize(resolution=WORK_AREA**0.5, mode="area", downsample_only=False),
133
+ Lambda(lambda x: torch.clamp(x, 0.0, 1.0)),
134
+ DivisibleCrop((16, 16)),
135
+ Normalize(0.5, 0.5),
136
+ Rearrange("t c h w -> c t h w"),
137
+ ]
138
+ )
139
+
140
+
141
+ def _duration(image, steps=1, seed=666, progress=None) -> int:
142
+ """Every request restores at the same ~3.7MP working resolution, so cost
143
+ tracks the step count and little else. Measured on a cold worker, then
144
+ given headroom; kept as small as honesty allows, because the request is
145
+ checked against the visitor's remaining quota and a smaller one also ranks
146
+ higher in the ZeroGPU queue."""
147
+ return int(min(120, 35 + 12 * int(steps)))
148
+
149
+
150
+ @spaces.GPU(duration=_duration)
151
+ @torch.no_grad()
152
+ def upscale_image(image, steps=1, seed=666, progress=gr.Progress(track_tqdm=True)):
153
+ if image is None:
154
+ raise gr.Error("Upload an image first.")
155
+
156
+ runner = _ensure_runner()
157
+
158
+ runner.config.diffusion.cfg.scale = 1.0
159
+ runner.config.diffusion.cfg.rescale = 0.0
160
+ runner.config.diffusion.timesteps.sampling.steps = int(steps)
161
+ runner.configure_diffusion()
162
+ set_seed(int(seed) % (2**32), same_across_ranks=True)
163
+
164
+ img = Image.open(image).convert("RGB") if isinstance(image, str) else image.convert("RGB")
165
+ tensor = T.ToTensor()(img).unsqueeze(0) # (t=1, c, h, w)
166
+
167
+ cond = _transform()(tensor.to("cuda"))
168
+ original = cond
169
+ latents = runner.vae_encode([cond])
170
+
171
+ text_embeds = {
172
+ "texts_pos": [torch.load(POS_EMB).to("cuda")],
173
+ "texts_neg": [torch.load(NEG_EMB).to("cuda")],
174
+ }
175
+
176
+ noise = [torch.randn_like(latent) for latent in latents]
177
+ aug_noise = [torch.randn_like(latent) for latent in latents]
178
+
179
+ def _add_noise(x, aug):
180
+ t = torch.tensor([1000.0], device="cuda") * 0.1
181
+ shape = torch.tensor(x.shape[1:], device="cuda")[None]
182
+ return runner.schedule.forward(x, aug, runner.timestep_transform(t, shape))
183
+
184
+ conditions = [
185
+ runner.get_condition(n, task="sr", latent_blur=_add_noise(latent, a))
186
+ for n, a, latent in zip(noise, aug_noise, latents)
187
+ ]
188
+
189
+ with torch.autocast("cuda", torch.bfloat16, enabled=True):
190
+ videos = runner.inference(
191
+ noises=noise, conditions=conditions, dit_offload=False, **text_embeds
192
+ )
193
+
194
+ sample = videos[0]
195
+ sample = (
196
+ rearrange(sample[:, None], "c t h w -> t c h w")
197
+ if sample.ndim == 3
198
+ else rearrange(sample, "c t h w -> t c h w")
199
+ )
200
+ reference = (
201
+ rearrange(original[:, None], "c t h w -> t c h w")
202
+ if original.ndim == 3
203
+ else rearrange(original, "c t h w -> t c h w")
204
+ )
205
+ if USE_COLOR_FIX:
206
+ sample = wavelet_reconstruction(sample.to("cpu"), reference[: sample.size(0)].to("cpu"))
207
+ else:
208
+ sample = sample.to("cpu")
209
+
210
+ sample = rearrange(sample, "t c h w -> t h w c")
211
+ sample = sample.clip(-1, 1).mul_(0.5).add_(0.5).mul_(255).round().to(torch.uint8).numpy()
212
+
213
+ del latents, conditions, videos
214
+ gc.collect()
215
+ torch.cuda.empty_cache()
216
+
217
+ return Image.fromarray(sample[0])
218
+
219
+
220
+ with gr.Blocks(css=UPSAMPLER_CSS, theme=UPSAMPLER_THEME) as demo:
221
+ gr.HTML(
222
+ header_html(
223
+ "SeedVR2 3B Image Upscaler",
224
+ "One-step diffusion restoration that rebuilds real detail in blurry, "
225
+ "compressed, and low-resolution photos.",
226
+ )
227
+ )
228
+
229
+ with gr.Row():
230
+ with gr.Column():
231
+ image_in = gr.Image(label="Image", type="filepath")
232
+ steps = gr.Slider(1, 4, value=1, step=1, label="Steps")
233
+ seed = gr.Number(label="Seed", value=666, precision=0)
234
+ run = gr.Button("Upscale Image", variant="primary")
235
+ with gr.Column():
236
+ image_out = gr.Image(label="Result", type="pil")
237
+
238
+ # No leading slash: gradio prefixes it, and "/upscale_image" here would
239
+ # publish the endpoint as "//upscale_image".
240
+ run.click(
241
+ upscale_image,
242
+ inputs=[image_in, steps, seed],
243
+ outputs=[image_out],
244
+ api_name="upscale_image",
245
+ )
246
+
247
+ gr.HTML(
248
+ footer_html(
249
+ "SeedVR2-3B is ByteDance's one-step diffusion model for image and video "
250
+ "restoration. It rebuilds genuine texture in photos that are blurry, "
251
+ "heavily compressed, or simply too small, restoring at high resolution "
252
+ "rather than smoothing detail away the way a conventional upscaler does.",
253
+ "https://upsampler.com/free-image-upscaler-no-signup",
254
+ "free image upscaler",
255
+ )
256
+ )
257
+
258
+ demo.launch(ssr_mode=False, show_error=True)
common/__init__.py ADDED
File without changes
common/cache.py ADDED
@@ -0,0 +1,47 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # // Copyright (c) 2025 Bytedance Ltd. and/or its affiliates
2
+ # //
3
+ # // Licensed under the Apache License, Version 2.0 (the "License");
4
+ # // you may not use this file except in compliance with the License.
5
+ # // You may obtain a copy of the License at
6
+ # //
7
+ # // http://www.apache.org/licenses/LICENSE-2.0
8
+ # //
9
+ # // Unless required by applicable law or agreed to in writing, software
10
+ # // distributed under the License is distributed on an "AS IS" BASIS,
11
+ # // WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # // See the License for the specific language governing permissions and
13
+ # // limitations under the License.
14
+
15
+ from typing import Callable
16
+
17
+
18
+ class Cache:
19
+ """Caching reusable args for faster inference"""
20
+
21
+ def __init__(self, disable=False, prefix="", cache=None):
22
+ self.cache = cache if cache is not None else {}
23
+ self.disable = disable
24
+ self.prefix = prefix
25
+
26
+ def __call__(self, key: str, fn: Callable):
27
+ if self.disable:
28
+ return fn()
29
+
30
+ key = self.prefix + key
31
+ try:
32
+ result = self.cache[key]
33
+ except KeyError:
34
+ result = fn()
35
+ self.cache[key] = result
36
+ return result
37
+
38
+ def namespace(self, namespace: str):
39
+ return Cache(
40
+ disable=self.disable,
41
+ prefix=self.prefix + namespace + ".",
42
+ cache=self.cache,
43
+ )
44
+
45
+ def get(self, key: str):
46
+ key = self.prefix + key
47
+ return self.cache[key]
common/config.py ADDED
@@ -0,0 +1,110 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # // Copyright (c) 2025 Bytedance Ltd. and/or its affiliates
2
+ # //
3
+ # // Licensed under the Apache License, Version 2.0 (the "License");
4
+ # // you may not use this file except in compliance with the License.
5
+ # // You may obtain a copy of the License at
6
+ # //
7
+ # // http://www.apache.org/licenses/LICENSE-2.0
8
+ # //
9
+ # // Unless required by applicable law or agreed to in writing, software
10
+ # // distributed under the License is distributed on an "AS IS" BASIS,
11
+ # // WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # // See the License for the specific language governing permissions and
13
+ # // limitations under the License.
14
+
15
+ """
16
+ Configuration utility functions
17
+ """
18
+
19
+ import importlib
20
+ from typing import Any, Callable, List, Union
21
+ from omegaconf import DictConfig, ListConfig, OmegaConf
22
+
23
+ OmegaConf.register_new_resolver("eval", eval)
24
+
25
+
26
+ def load_config(path: str, argv: List[str] = None) -> Union[DictConfig, ListConfig]:
27
+ """
28
+ Load a configuration. Will resolve inheritance.
29
+ """
30
+ config = OmegaConf.load(path)
31
+ if argv is not None:
32
+ config_argv = OmegaConf.from_dotlist(argv)
33
+ config = OmegaConf.merge(config, config_argv)
34
+ config = resolve_recursive(config, resolve_inheritance)
35
+ return config
36
+
37
+
38
+ def resolve_recursive(
39
+ config: Any,
40
+ resolver: Callable[[Union[DictConfig, ListConfig]], Union[DictConfig, ListConfig]],
41
+ ) -> Any:
42
+ config = resolver(config)
43
+ if isinstance(config, DictConfig):
44
+ for k in config.keys():
45
+ v = config.get(k)
46
+ if isinstance(v, (DictConfig, ListConfig)):
47
+ config[k] = resolve_recursive(v, resolver)
48
+ if isinstance(config, ListConfig):
49
+ for i in range(len(config)):
50
+ v = config.get(i)
51
+ if isinstance(v, (DictConfig, ListConfig)):
52
+ config[i] = resolve_recursive(v, resolver)
53
+ return config
54
+
55
+
56
+ def resolve_inheritance(config: Union[DictConfig, ListConfig]) -> Any:
57
+ """
58
+ Recursively resolve inheritance if the config contains:
59
+ __inherit__: path/to/parent.yaml or a ListConfig of such paths.
60
+ """
61
+ if isinstance(config, DictConfig):
62
+ inherit = config.pop("__inherit__", None)
63
+
64
+ if inherit:
65
+ inherit_list = inherit if isinstance(inherit, ListConfig) else [inherit]
66
+
67
+ parent_config = None
68
+ for parent_path in inherit_list:
69
+ assert isinstance(parent_path, str)
70
+ parent_config = (
71
+ load_config(parent_path)
72
+ if parent_config is None
73
+ else OmegaConf.merge(parent_config, load_config(parent_path))
74
+ )
75
+
76
+ if len(config.keys()) > 0:
77
+ config = OmegaConf.merge(parent_config, config)
78
+ else:
79
+ config = parent_config
80
+ return config
81
+
82
+
83
+ def import_item(path: str, name: str) -> Any:
84
+ """
85
+ Import a python item. Example: import_item("path.to.file", "MyClass") -> MyClass
86
+ """
87
+ return getattr(importlib.import_module(path), name)
88
+
89
+
90
+ def create_object(config: DictConfig) -> Any:
91
+ """
92
+ Create an object from config.
93
+ The config is expected to contains the following:
94
+ __object__:
95
+ path: path.to.module
96
+ name: MyClass
97
+ args: as_config | as_params (default to as_config)
98
+ """
99
+ item = import_item(
100
+ path=config.__object__.path,
101
+ name=config.__object__.name,
102
+ )
103
+ args = config.__object__.get("args", "as_config")
104
+ if args == "as_config":
105
+ return item(config)
106
+ if args == "as_params":
107
+ config = OmegaConf.to_object(config)
108
+ config.pop("__object__")
109
+ return item(**config)
110
+ raise NotImplementedError(f"Unknown args type: {args}")
common/decorators.py ADDED
@@ -0,0 +1,147 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # // Copyright (c) 2025 Bytedance Ltd. and/or its affiliates
2
+ # //
3
+ # // Licensed under the Apache License, Version 2.0 (the "License");
4
+ # // you may not use this file except in compliance with the License.
5
+ # // You may obtain a copy of the License at
6
+ # //
7
+ # // http://www.apache.org/licenses/LICENSE-2.0
8
+ # //
9
+ # // Unless required by applicable law or agreed to in writing, software
10
+ # // distributed under the License is distributed on an "AS IS" BASIS,
11
+ # // WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # // See the License for the specific language governing permissions and
13
+ # // limitations under the License.
14
+
15
+ """
16
+ Decorators.
17
+ """
18
+
19
+ import functools
20
+ import threading
21
+ import time
22
+ from typing import Callable
23
+ import torch
24
+
25
+ from common.distributed import barrier_if_distributed, get_global_rank, get_local_rank
26
+ from common.logger import get_logger
27
+
28
+ logger = get_logger(__name__)
29
+
30
+
31
+ def log_on_entry(func: Callable) -> Callable:
32
+ """
33
+ Functions with this decorator will log the function name at entry.
34
+ When using multiple decorators, this must be applied innermost to properly capture the name.
35
+ """
36
+
37
+ def log_on_entry_wrapper(*args, **kwargs):
38
+ logger.info(f"Entering {func.__name__}")
39
+ return func(*args, **kwargs)
40
+
41
+ return log_on_entry_wrapper
42
+
43
+
44
+ def barrier_on_entry(func: Callable) -> Callable:
45
+ """
46
+ Functions with this decorator will start executing when all ranks are ready to enter.
47
+ """
48
+
49
+ def barrier_on_entry_wrapper(*args, **kwargs):
50
+ barrier_if_distributed()
51
+ return func(*args, **kwargs)
52
+
53
+ return barrier_on_entry_wrapper
54
+
55
+
56
+ def _conditional_execute_wrapper_factory(execute: bool, func: Callable) -> Callable:
57
+ """
58
+ Helper function for local_rank_zero_only and global_rank_zero_only.
59
+ """
60
+
61
+ def conditional_execute_wrapper(*args, **kwargs):
62
+ # Only execute if needed.
63
+ result = func(*args, **kwargs) if execute else None
64
+ # All GPUs must wait.
65
+ barrier_if_distributed()
66
+ # Return results.
67
+ return result
68
+
69
+ return conditional_execute_wrapper
70
+
71
+
72
+ def _asserted_wrapper_factory(condition: bool, func: Callable, err_msg: str = "") -> Callable:
73
+ """
74
+ Helper function for some functions with special constraints,
75
+ especially functions called by other global_rank_zero_only / local_rank_zero_only ones,
76
+ in case they are wrongly invoked in other scenarios.
77
+ """
78
+
79
+ def asserted_execute_wrapper(*args, **kwargs):
80
+ assert condition, err_msg
81
+ result = func(*args, **kwargs)
82
+ return result
83
+
84
+ return asserted_execute_wrapper
85
+
86
+
87
+ def local_rank_zero_only(func: Callable) -> Callable:
88
+ """
89
+ Functions with this decorator will only execute on local rank zero.
90
+ """
91
+ return _conditional_execute_wrapper_factory(get_local_rank() == 0, func)
92
+
93
+
94
+ def global_rank_zero_only(func: Callable) -> Callable:
95
+ """
96
+ Functions with this decorator will only execute on global rank zero.
97
+ """
98
+ return _conditional_execute_wrapper_factory(get_global_rank() == 0, func)
99
+
100
+
101
+ def assert_only_global_rank_zero(func: Callable) -> Callable:
102
+ """
103
+ Functions with this decorator are only accessible to processes with global rank zero.
104
+ """
105
+ return _asserted_wrapper_factory(
106
+ get_global_rank() == 0, func, err_msg="Not accessible to processes with global_rank != 0"
107
+ )
108
+
109
+
110
+ def assert_only_local_rank_zero(func: Callable) -> Callable:
111
+ """
112
+ Functions with this decorator are only accessible to processes with local rank zero.
113
+ """
114
+ return _asserted_wrapper_factory(
115
+ get_local_rank() == 0, func, err_msg="Not accessible to processes with local_rank != 0"
116
+ )
117
+
118
+
119
+ def new_thread(func: Callable) -> Callable:
120
+ """
121
+ Functions with this decorator will run in a new thread.
122
+ The function will return the thread, which can be joined to wait for completion.
123
+ """
124
+
125
+ def new_thread_wrapper(*args, **kwargs):
126
+ thread = threading.Thread(target=func, args=args, kwargs=kwargs)
127
+ thread.start()
128
+ return thread
129
+
130
+ return new_thread_wrapper
131
+
132
+
133
+ def log_runtime(func: Callable) -> Callable:
134
+ """
135
+ Functions with this decorator will logging the runtime.
136
+ """
137
+
138
+ @functools.wraps(func)
139
+ def wrapped(*args, **kwargs):
140
+ torch.distributed.barrier()
141
+ start = time.perf_counter()
142
+ result = func(*args, **kwargs)
143
+ torch.distributed.barrier()
144
+ logger.info(f"Completed {func.__name__} in {time.perf_counter() - start:.3f} seconds.")
145
+ return result
146
+
147
+ return wrapped
common/diffusion/__init__.py ADDED
@@ -0,0 +1,56 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # // Copyright (c) 2025 Bytedance Ltd. and/or its affiliates
2
+ # //
3
+ # // Licensed under the Apache License, Version 2.0 (the "License");
4
+ # // you may not use this file except in compliance with the License.
5
+ # // You may obtain a copy of the License at
6
+ # //
7
+ # // http://www.apache.org/licenses/LICENSE-2.0
8
+ # //
9
+ # // Unless required by applicable law or agreed to in writing, software
10
+ # // distributed under the License is distributed on an "AS IS" BASIS,
11
+ # // WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # // See the License for the specific language governing permissions and
13
+ # // limitations under the License.
14
+
15
+ """
16
+ Diffusion package.
17
+ """
18
+
19
+ from .config import (
20
+ create_sampler_from_config,
21
+ create_sampling_timesteps_from_config,
22
+ create_schedule_from_config,
23
+ )
24
+ from .samplers.base import Sampler
25
+ from .samplers.euler import EulerSampler
26
+ from .schedules.base import Schedule
27
+ from .schedules.lerp import LinearInterpolationSchedule
28
+ from .timesteps.base import SamplingTimesteps, Timesteps
29
+ from .timesteps.sampling.trailing import UniformTrailingSamplingTimesteps
30
+ from .types import PredictionType, SamplingDirection
31
+ from .utils import classifier_free_guidance, classifier_free_guidance_dispatcher, expand_dims
32
+
33
+ __all__ = [
34
+ # Configs
35
+ "create_sampler_from_config",
36
+ "create_sampling_timesteps_from_config",
37
+ "create_schedule_from_config",
38
+ # Schedules
39
+ "Schedule",
40
+ "DiscreteVariancePreservingSchedule",
41
+ "LinearInterpolationSchedule",
42
+ # Samplers
43
+ "Sampler",
44
+ "EulerSampler",
45
+ # Timesteps
46
+ "Timesteps",
47
+ "SamplingTimesteps",
48
+ # Types
49
+ "PredictionType",
50
+ "SamplingDirection",
51
+ "UniformTrailingSamplingTimesteps",
52
+ # Utils
53
+ "classifier_free_guidance",
54
+ "classifier_free_guidance_dispatcher",
55
+ "expand_dims",
56
+ ]
common/diffusion/config.py ADDED
@@ -0,0 +1,74 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # // Copyright (c) 2025 Bytedance Ltd. and/or its affiliates
2
+ # //
3
+ # // Licensed under the Apache License, Version 2.0 (the "License");
4
+ # // you may not use this file except in compliance with the License.
5
+ # // You may obtain a copy of the License at
6
+ # //
7
+ # // http://www.apache.org/licenses/LICENSE-2.0
8
+ # //
9
+ # // Unless required by applicable law or agreed to in writing, software
10
+ # // distributed under the License is distributed on an "AS IS" BASIS,
11
+ # // WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # // See the License for the specific language governing permissions and
13
+ # // limitations under the License.
14
+
15
+ """
16
+ Utility functions for creating schedules and samplers from config.
17
+ """
18
+
19
+ import torch
20
+ from omegaconf import DictConfig
21
+
22
+ from .samplers.base import Sampler
23
+ from .samplers.euler import EulerSampler
24
+ from .schedules.base import Schedule
25
+ from .schedules.lerp import LinearInterpolationSchedule
26
+ from .timesteps.base import SamplingTimesteps
27
+ from .timesteps.sampling.trailing import UniformTrailingSamplingTimesteps
28
+
29
+
30
+ def create_schedule_from_config(
31
+ config: DictConfig,
32
+ device: torch.device,
33
+ dtype: torch.dtype = torch.float32,
34
+ ) -> Schedule:
35
+ """
36
+ Create a schedule from configuration.
37
+ """
38
+ if config.type == "lerp":
39
+ return LinearInterpolationSchedule(T=config.get("T", 1.0))
40
+
41
+ raise NotImplementedError
42
+
43
+
44
+ def create_sampler_from_config(
45
+ config: DictConfig,
46
+ schedule: Schedule,
47
+ timesteps: SamplingTimesteps,
48
+ ) -> Sampler:
49
+ """
50
+ Create a sampler from configuration.
51
+ """
52
+ if config.type == "euler":
53
+ return EulerSampler(
54
+ schedule=schedule,
55
+ timesteps=timesteps,
56
+ prediction_type=config.prediction_type,
57
+ )
58
+ raise NotImplementedError
59
+
60
+
61
+ def create_sampling_timesteps_from_config(
62
+ config: DictConfig,
63
+ schedule: Schedule,
64
+ device: torch.device,
65
+ dtype: torch.dtype = torch.float32,
66
+ ) -> SamplingTimesteps:
67
+ if config.type == "uniform_trailing":
68
+ return UniformTrailingSamplingTimesteps(
69
+ T=schedule.T,
70
+ steps=config.steps,
71
+ shift=config.get("shift", 1.0),
72
+ device=device,
73
+ )
74
+ raise NotImplementedError
common/diffusion/samplers/base.py ADDED
@@ -0,0 +1,108 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # // Copyright (c) 2025 Bytedance Ltd. and/or its affiliates
2
+ # //
3
+ # // Licensed under the Apache License, Version 2.0 (the "License");
4
+ # // you may not use this file except in compliance with the License.
5
+ # // You may obtain a copy of the License at
6
+ # //
7
+ # // http://www.apache.org/licenses/LICENSE-2.0
8
+ # //
9
+ # // Unless required by applicable law or agreed to in writing, software
10
+ # // distributed under the License is distributed on an "AS IS" BASIS,
11
+ # // WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # // See the License for the specific language governing permissions and
13
+ # // limitations under the License.
14
+
15
+ """
16
+ Sampler base class.
17
+ """
18
+
19
+ from abc import ABC, abstractmethod
20
+ from dataclasses import dataclass
21
+ from typing import Callable
22
+ import torch
23
+ from tqdm import tqdm
24
+
25
+ from ..schedules.base import Schedule
26
+ from ..timesteps.base import SamplingTimesteps
27
+ from ..types import PredictionType, SamplingDirection
28
+ from ..utils import assert_schedule_timesteps_compatible
29
+
30
+
31
+ @dataclass
32
+ class SamplerModelArgs:
33
+ x_t: torch.Tensor
34
+ t: torch.Tensor
35
+ i: int
36
+
37
+
38
+ class Sampler(ABC):
39
+ """
40
+ Samplers are ODE/SDE solvers.
41
+ """
42
+
43
+ def __init__(
44
+ self,
45
+ schedule: Schedule,
46
+ timesteps: SamplingTimesteps,
47
+ prediction_type: PredictionType,
48
+ return_endpoint: bool = True,
49
+ ):
50
+ assert_schedule_timesteps_compatible(
51
+ schedule=schedule,
52
+ timesteps=timesteps,
53
+ )
54
+ self.schedule = schedule
55
+ self.timesteps = timesteps
56
+ self.prediction_type = prediction_type
57
+ self.return_endpoint = return_endpoint
58
+
59
+ @abstractmethod
60
+ def sample(
61
+ self,
62
+ x: torch.Tensor,
63
+ f: Callable[[SamplerModelArgs], torch.Tensor],
64
+ ) -> torch.Tensor:
65
+ """
66
+ Generate a new sample given the the intial sample x and score function f.
67
+ """
68
+
69
+ def get_next_timestep(
70
+ self,
71
+ t: torch.Tensor,
72
+ ) -> torch.Tensor:
73
+ """
74
+ Get the next sample timestep.
75
+ Support multiple different timesteps t in a batch.
76
+ If no more steps, return out of bound value -1 or T+1.
77
+ """
78
+ T = self.timesteps.T
79
+ steps = len(self.timesteps)
80
+ curr_idx = self.timesteps.index(t)
81
+ next_idx = curr_idx + 1
82
+ bound = -1 if self.timesteps.direction == SamplingDirection.backward else T + 1
83
+
84
+ s = self.timesteps[next_idx.clamp_max(steps - 1)]
85
+ s = s.where(next_idx < steps, bound)
86
+ return s
87
+
88
+ def get_endpoint(
89
+ self,
90
+ pred: torch.Tensor,
91
+ x_t: torch.Tensor,
92
+ t: torch.Tensor,
93
+ ) -> torch.Tensor:
94
+ """
95
+ Get to the endpoint of the probability flow.
96
+ """
97
+ x_0, x_T = self.schedule.convert_from_pred(pred, self.prediction_type, x_t, t)
98
+ return x_0 if self.timesteps.direction == SamplingDirection.backward else x_T
99
+
100
+ def get_progress_bar(self):
101
+ """
102
+ Get progress bar for sampling.
103
+ """
104
+ return tqdm(
105
+ iterable=range(len(self.timesteps) - (0 if self.return_endpoint else 1)),
106
+ dynamic_ncols=True,
107
+ desc=self.__class__.__name__,
108
+ )
common/diffusion/samplers/euler.py ADDED
@@ -0,0 +1,89 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # // Copyright (c) 2025 Bytedance Ltd. and/or its affiliates
2
+ # //
3
+ # // Licensed under the Apache License, Version 2.0 (the "License");
4
+ # // you may not use this file except in compliance with the License.
5
+ # // You may obtain a copy of the License at
6
+ # //
7
+ # // http://www.apache.org/licenses/LICENSE-2.0
8
+ # //
9
+ # // Unless required by applicable law or agreed to in writing, software
10
+ # // distributed under the License is distributed on an "AS IS" BASIS,
11
+ # // WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # // See the License for the specific language governing permissions and
13
+ # // limitations under the License.
14
+
15
+
16
+ """
17
+ Euler ODE solver.
18
+ """
19
+
20
+ from typing import Callable
21
+ import torch
22
+ from einops import rearrange
23
+ from torch.nn import functional as F
24
+
25
+ from models.dit_v2 import na
26
+
27
+ from ..types import PredictionType
28
+ from ..utils import expand_dims
29
+ from .base import Sampler, SamplerModelArgs
30
+
31
+
32
+ class EulerSampler(Sampler):
33
+ """
34
+ The Euler method is the simplest ODE solver.
35
+ <https://en.wikipedia.org/wiki/Euler_method>
36
+ """
37
+
38
+ def sample(
39
+ self,
40
+ x: torch.Tensor,
41
+ f: Callable[[SamplerModelArgs], torch.Tensor],
42
+ ) -> torch.Tensor:
43
+ timesteps = self.timesteps.timesteps
44
+ progress = self.get_progress_bar()
45
+ i = 0
46
+ for t, s in zip(timesteps[:-1], timesteps[1:]):
47
+ pred = f(SamplerModelArgs(x, t, i))
48
+ x = self.step_to(pred, x, t, s)
49
+ i += 1
50
+ progress.update()
51
+
52
+ if self.return_endpoint:
53
+ t = timesteps[-1]
54
+ pred = f(SamplerModelArgs(x, t, i))
55
+ x = self.get_endpoint(pred, x, t)
56
+ progress.update()
57
+ return x
58
+
59
+ def step(
60
+ self,
61
+ pred: torch.Tensor,
62
+ x_t: torch.Tensor,
63
+ t: torch.Tensor,
64
+ ) -> torch.Tensor:
65
+ """
66
+ Step to the next timestep.
67
+ """
68
+ return self.step_to(pred, x_t, t, self.get_next_timestep(t))
69
+
70
+ def step_to(
71
+ self,
72
+ pred: torch.Tensor,
73
+ x_t: torch.Tensor,
74
+ t: torch.Tensor,
75
+ s: torch.Tensor,
76
+ ) -> torch.Tensor:
77
+ """
78
+ Steps from x_t at timestep t to x_s at timestep s. Returns x_s.
79
+ """
80
+ t = expand_dims(t, x_t.ndim)
81
+ s = expand_dims(s, x_t.ndim)
82
+ T = self.schedule.T
83
+ # Step from x_t to x_s.
84
+ pred_x_0, pred_x_T = self.schedule.convert_from_pred(pred, self.prediction_type, x_t, t)
85
+ pred_x_s = self.schedule.forward(pred_x_0, pred_x_T, s.clamp(0, T))
86
+ # Clamp x_s to x_0 and x_T if s is out of bound.
87
+ pred_x_s = pred_x_s.where(s >= 0, pred_x_0)
88
+ pred_x_s = pred_x_s.where(s <= T, pred_x_T)
89
+ return pred_x_s
common/diffusion/schedules/base.py ADDED
@@ -0,0 +1,131 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # // Copyright (c) 2025 Bytedance Ltd. and/or its affiliates
2
+ # //
3
+ # // Licensed under the Apache License, Version 2.0 (the "License");
4
+ # // you may not use this file except in compliance with the License.
5
+ # // You may obtain a copy of the License at
6
+ # //
7
+ # // http://www.apache.org/licenses/LICENSE-2.0
8
+ # //
9
+ # // Unless required by applicable law or agreed to in writing, software
10
+ # // distributed under the License is distributed on an "AS IS" BASIS,
11
+ # // WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # // See the License for the specific language governing permissions and
13
+ # // limitations under the License.
14
+
15
+ """
16
+ Schedule base class.
17
+ """
18
+
19
+ from abc import ABC, abstractmethod, abstractproperty
20
+ from typing import Tuple, Union
21
+ import torch
22
+
23
+ from ..types import PredictionType
24
+ from ..utils import expand_dims
25
+
26
+
27
+ class Schedule(ABC):
28
+ """
29
+ Diffusion schedules are uniquely defined by T, A, B:
30
+
31
+ x_t = A(t) * x_0 + B(t) * x_T, where t in [0, T]
32
+
33
+ Schedules can be continuous or discrete.
34
+ """
35
+
36
+ @abstractproperty
37
+ def T(self) -> Union[int, float]:
38
+ """
39
+ Maximum timestep inclusive.
40
+ Schedule is continuous if float, discrete if int.
41
+ """
42
+
43
+ @abstractmethod
44
+ def A(self, t: torch.Tensor) -> torch.Tensor:
45
+ """
46
+ Interpolation coefficient A.
47
+ Returns tensor with the same shape as t.
48
+ """
49
+
50
+ @abstractmethod
51
+ def B(self, t: torch.Tensor) -> torch.Tensor:
52
+ """
53
+ Interpolation coefficient B.
54
+ Returns tensor with the same shape as t.
55
+ """
56
+
57
+ # ----------------------------------------------------
58
+
59
+ def snr(self, t: torch.Tensor) -> torch.Tensor:
60
+ """
61
+ Signal to noise ratio.
62
+ Returns tensor with the same shape as t.
63
+ """
64
+ return (self.A(t) ** 2) / (self.B(t) ** 2)
65
+
66
+ def isnr(self, snr: torch.Tensor) -> torch.Tensor:
67
+ """
68
+ Inverse signal to noise ratio.
69
+ Returns tensor with the same shape as snr.
70
+ Subclass may implement.
71
+ """
72
+ raise NotImplementedError
73
+
74
+ # ----------------------------------------------------
75
+
76
+ def is_continuous(self) -> bool:
77
+ """
78
+ Whether the schedule is continuous.
79
+ """
80
+ return isinstance(self.T, float)
81
+
82
+ def forward(self, x_0: torch.Tensor, x_T: torch.Tensor, t: torch.Tensor) -> torch.Tensor:
83
+ """
84
+ Diffusion forward function.
85
+ """
86
+ t = expand_dims(t, x_0.ndim)
87
+ return self.A(t) * x_0 + self.B(t) * x_T
88
+
89
+ def convert_from_pred(
90
+ self, pred: torch.Tensor, pred_type: PredictionType, x_t: torch.Tensor, t: torch.Tensor
91
+ ) -> Tuple[torch.Tensor, torch.Tensor]:
92
+ """
93
+ Convert from prediction. Return predicted x_0 and x_T.
94
+ """
95
+ t = expand_dims(t, x_t.ndim)
96
+ A_t = self.A(t)
97
+ B_t = self.B(t)
98
+
99
+ if pred_type == PredictionType.x_T:
100
+ pred_x_T = pred
101
+ pred_x_0 = (x_t - B_t * pred_x_T) / A_t
102
+ elif pred_type == PredictionType.x_0:
103
+ pred_x_0 = pred
104
+ pred_x_T = (x_t - A_t * pred_x_0) / B_t
105
+ elif pred_type == PredictionType.v_cos:
106
+ pred_x_0 = A_t * x_t - B_t * pred
107
+ pred_x_T = A_t * pred + B_t * x_t
108
+ elif pred_type == PredictionType.v_lerp:
109
+ pred_x_0 = (x_t - B_t * pred) / (A_t + B_t)
110
+ pred_x_T = (x_t + A_t * pred) / (A_t + B_t)
111
+ else:
112
+ raise NotImplementedError
113
+
114
+ return pred_x_0, pred_x_T
115
+
116
+ def convert_to_pred(
117
+ self, x_0: torch.Tensor, x_T: torch.Tensor, t: torch.Tensor, pred_type: PredictionType
118
+ ) -> torch.FloatTensor:
119
+ """
120
+ Convert to prediction target given x_0 and x_T.
121
+ """
122
+ if pred_type == PredictionType.x_T:
123
+ return x_T
124
+ if pred_type == PredictionType.x_0:
125
+ return x_0
126
+ if pred_type == PredictionType.v_cos:
127
+ t = expand_dims(t, x_0.ndim)
128
+ return self.A(t) * x_T - self.B(t) * x_0
129
+ if pred_type == PredictionType.v_lerp:
130
+ return x_T - x_0
131
+ raise NotImplementedError
common/diffusion/schedules/lerp.py ADDED
@@ -0,0 +1,55 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # // Copyright (c) 2025 Bytedance Ltd. and/or its affiliates
2
+ # //
3
+ # // Licensed under the Apache License, Version 2.0 (the "License");
4
+ # // you may not use this file except in compliance with the License.
5
+ # // You may obtain a copy of the License at
6
+ # //
7
+ # // http://www.apache.org/licenses/LICENSE-2.0
8
+ # //
9
+ # // Unless required by applicable law or agreed to in writing, software
10
+ # // distributed under the License is distributed on an "AS IS" BASIS,
11
+ # // WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # // See the License for the specific language governing permissions and
13
+ # // limitations under the License.
14
+
15
+ """
16
+ Linear interpolation schedule (lerp).
17
+ """
18
+
19
+ from typing import Union
20
+ import torch
21
+
22
+ from .base import Schedule
23
+
24
+
25
+ class LinearInterpolationSchedule(Schedule):
26
+ """
27
+ Linear interpolation schedule (lerp) is proposed by flow matching and rectified flow.
28
+ It leads to straighter probability flow theoretically. It is also used by Stable Diffusion 3.
29
+ <https://arxiv.org/abs/2209.03003>
30
+ <https://arxiv.org/abs/2210.02747>
31
+
32
+ x_t = (1 - t) * x_0 + t * x_T
33
+
34
+ Can be either continuous or discrete.
35
+ """
36
+
37
+ def __init__(self, T: Union[int, float] = 1.0):
38
+ self._T = T
39
+
40
+ @property
41
+ def T(self) -> Union[int, float]:
42
+ return self._T
43
+
44
+ def A(self, t: torch.Tensor) -> torch.Tensor:
45
+ return 1 - (t / self.T)
46
+
47
+ def B(self, t: torch.Tensor) -> torch.Tensor:
48
+ return t / self.T
49
+
50
+ # ----------------------------------------------------
51
+
52
+ def isnr(self, snr: torch.Tensor) -> torch.Tensor:
53
+ t = self.T / (1 + snr**0.5)
54
+ t = t if self.is_continuous() else t.round().int()
55
+ return t
common/diffusion/timesteps/base.py ADDED
@@ -0,0 +1,72 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from abc import ABC, abstractmethod
2
+ from typing import Sequence, Union
3
+ import torch
4
+
5
+ from ..types import SamplingDirection
6
+
7
+
8
+ class Timesteps(ABC):
9
+ """
10
+ Timesteps base class.
11
+ """
12
+
13
+ def __init__(self, T: Union[int, float]):
14
+ assert T > 0
15
+ self._T = T
16
+
17
+ @property
18
+ def T(self) -> Union[int, float]:
19
+ """
20
+ Maximum timestep inclusive.
21
+ int if discrete, float if continuous.
22
+ """
23
+ return self._T
24
+
25
+ def is_continuous(self) -> bool:
26
+ """
27
+ Whether the schedule is continuous.
28
+ """
29
+ return isinstance(self.T, float)
30
+
31
+
32
+ class SamplingTimesteps(Timesteps):
33
+ """
34
+ Sampling timesteps.
35
+ It defines the discretization of sampling steps.
36
+ """
37
+
38
+ def __init__(
39
+ self,
40
+ T: Union[int, float],
41
+ timesteps: torch.Tensor,
42
+ direction: SamplingDirection,
43
+ ):
44
+ assert timesteps.ndim == 1
45
+ super().__init__(T)
46
+ self.timesteps = timesteps
47
+ self.direction = direction
48
+
49
+ def __len__(self) -> int:
50
+ """
51
+ Number of sampling steps.
52
+ """
53
+ return len(self.timesteps)
54
+
55
+ def __getitem__(self, idx: Union[int, torch.IntTensor]) -> torch.Tensor:
56
+ """
57
+ The timestep at the sampling step.
58
+ Returns a scalar tensor if idx is int,
59
+ or tensor of the same size if idx is a tensor.
60
+ """
61
+ return self.timesteps[idx]
62
+
63
+ def index(self, t: torch.Tensor) -> torch.Tensor:
64
+ """
65
+ Find index by t.
66
+ Return index of the same shape as t.
67
+ Index is -1 if t not found in timesteps.
68
+ """
69
+ i, j = t.reshape(-1, 1).eq(self.timesteps).nonzero(as_tuple=True)
70
+ idx = torch.full_like(t, fill_value=-1, dtype=torch.int)
71
+ idx.view(-1)[i] = j.int()
72
+ return idx
common/diffusion/timesteps/sampling/trailing.py ADDED
@@ -0,0 +1,49 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # // Copyright (c) 2025 Bytedance Ltd. and/or its affiliates
2
+ # //
3
+ # // Licensed under the Apache License, Version 2.0 (the "License");
4
+ # // you may not use this file except in compliance with the License.
5
+ # // You may obtain a copy of the License at
6
+ # //
7
+ # // http://www.apache.org/licenses/LICENSE-2.0
8
+ # //
9
+ # // Unless required by applicable law or agreed to in writing, software
10
+ # // distributed under the License is distributed on an "AS IS" BASIS,
11
+ # // WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # // See the License for the specific language governing permissions and
13
+ # // limitations under the License.
14
+
15
+ import torch
16
+
17
+ from ...types import SamplingDirection
18
+ from ..base import SamplingTimesteps
19
+
20
+
21
+ class UniformTrailingSamplingTimesteps(SamplingTimesteps):
22
+ """
23
+ Uniform trailing sampling timesteps.
24
+ Defined in (https://arxiv.org/abs/2305.08891)
25
+
26
+ Shift is proposed in SD3 for RF schedule.
27
+ Defined in (https://arxiv.org/pdf/2403.03206) eq.23
28
+ """
29
+
30
+ def __init__(
31
+ self,
32
+ T: int,
33
+ steps: int,
34
+ shift: float = 1.0,
35
+ device: torch.device = "cpu",
36
+ ):
37
+ # Create trailing timesteps.
38
+ timesteps = torch.arange(1.0, 0.0, -1.0 / steps, device=device)
39
+
40
+ # Shift timesteps.
41
+ timesteps = shift * timesteps / (1 + (shift - 1) * timesteps)
42
+
43
+ # Scale to T range.
44
+ if isinstance(T, float):
45
+ timesteps = timesteps * T
46
+ else:
47
+ timesteps = timesteps.mul(T + 1).sub(1).round().int()
48
+
49
+ super().__init__(T=T, timesteps=timesteps, direction=SamplingDirection.backward)
common/diffusion/types.py ADDED
@@ -0,0 +1,59 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # // Copyright (c) 2025 Bytedance Ltd. and/or its affiliates
2
+ # //
3
+ # // Licensed under the Apache License, Version 2.0 (the "License");
4
+ # // you may not use this file except in compliance with the License.
5
+ # // You may obtain a copy of the License at
6
+ # //
7
+ # // http://www.apache.org/licenses/LICENSE-2.0
8
+ # //
9
+ # // Unless required by applicable law or agreed to in writing, software
10
+ # // distributed under the License is distributed on an "AS IS" BASIS,
11
+ # // WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # // See the License for the specific language governing permissions and
13
+ # // limitations under the License.
14
+
15
+ """
16
+ Type definitions.
17
+ """
18
+
19
+ from enum import Enum
20
+
21
+
22
+ class PredictionType(str, Enum):
23
+ """
24
+ x_0:
25
+ Predict data sample.
26
+ x_T:
27
+ Predict noise sample.
28
+ Proposed by DDPM (https://arxiv.org/abs/2006.11239)
29
+ Proved problematic by zsnr paper (https://arxiv.org/abs/2305.08891)
30
+ v_cos:
31
+ Predict velocity dx/dt based on the cosine schedule (A_t * x_T - B_t * x_0).
32
+ Proposed by progressive distillation (https://arxiv.org/abs/2202.00512)
33
+ v_lerp:
34
+ Predict velocity dx/dt based on the lerp schedule (x_T - x_0).
35
+ Proposed by rectified flow (https://arxiv.org/abs/2209.03003)
36
+ """
37
+
38
+ x_0 = "x_0"
39
+ x_T = "x_T"
40
+ v_cos = "v_cos"
41
+ v_lerp = "v_lerp"
42
+
43
+
44
+ class SamplingDirection(str, Enum):
45
+ """
46
+ backward: Sample from x_T to x_0 for data generation.
47
+ forward: Sample from x_0 to x_T for noise inversion.
48
+ """
49
+
50
+ backward = "backward"
51
+ forward = "forward"
52
+
53
+ @staticmethod
54
+ def reverse(direction):
55
+ if direction == SamplingDirection.backward:
56
+ return SamplingDirection.forward
57
+ if direction == SamplingDirection.forward:
58
+ return SamplingDirection.backward
59
+ raise NotImplementedError
common/diffusion/utils.py ADDED
@@ -0,0 +1,84 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # // Copyright (c) 2025 Bytedance Ltd. and/or its affiliates
2
+ # //
3
+ # // Licensed under the Apache License, Version 2.0 (the "License");
4
+ # // you may not use this file except in compliance with the License.
5
+ # // You may obtain a copy of the License at
6
+ # //
7
+ # // http://www.apache.org/licenses/LICENSE-2.0
8
+ # //
9
+ # // Unless required by applicable law or agreed to in writing, software
10
+ # // distributed under the License is distributed on an "AS IS" BASIS,
11
+ # // WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # // See the License for the specific language governing permissions and
13
+ # // limitations under the License.
14
+
15
+ """
16
+ Utility functions.
17
+ """
18
+
19
+ from typing import Callable
20
+ import torch
21
+
22
+
23
+ def expand_dims(tensor: torch.Tensor, ndim: int):
24
+ """
25
+ Expand tensor to target ndim. New dims are added to the right.
26
+ For example, if the tensor shape was (8,), target ndim is 4, return (8, 1, 1, 1).
27
+ """
28
+ shape = tensor.shape + (1,) * (ndim - tensor.ndim)
29
+ return tensor.reshape(shape)
30
+
31
+
32
+ def assert_schedule_timesteps_compatible(schedule, timesteps):
33
+ """
34
+ Check if schedule and timesteps are compatible.
35
+ """
36
+ if schedule.T != timesteps.T:
37
+ raise ValueError("Schedule and timesteps must have the same T.")
38
+ if schedule.is_continuous() != timesteps.is_continuous():
39
+ raise ValueError("Schedule and timesteps must have the same continuity.")
40
+
41
+
42
+ def classifier_free_guidance(
43
+ pos: torch.Tensor,
44
+ neg: torch.Tensor,
45
+ scale: float,
46
+ rescale: float = 0.0,
47
+ ):
48
+ """
49
+ Apply classifier-free guidance.
50
+ """
51
+ # Classifier-free guidance (https://arxiv.org/abs/2207.12598)
52
+ cfg = neg + scale * (pos - neg)
53
+
54
+ # Classifier-free guidance rescale (https://arxiv.org/pdf/2305.08891.pdf)
55
+ if rescale != 0.0:
56
+ pos_std = pos.std(dim=list(range(1, pos.ndim)), keepdim=True)
57
+ cfg_std = cfg.std(dim=list(range(1, cfg.ndim)), keepdim=True)
58
+ factor = pos_std / cfg_std
59
+ factor = rescale * factor + (1 - rescale)
60
+ cfg *= factor
61
+
62
+ return cfg
63
+
64
+
65
+ def classifier_free_guidance_dispatcher(
66
+ pos: Callable,
67
+ neg: Callable,
68
+ scale: float,
69
+ rescale: float = 0.0,
70
+ ):
71
+ """
72
+ Optionally execute models depending on classifer-free guidance scale.
73
+ """
74
+ # If scale is 1, no need to execute neg model.
75
+ if scale == 1.0:
76
+ return pos()
77
+
78
+ # Otherwise, execute both pos nad neg models and apply cfg.
79
+ return classifier_free_guidance(
80
+ pos=pos(),
81
+ neg=neg(),
82
+ scale=scale,
83
+ rescale=rescale,
84
+ )
common/distributed/__init__.py ADDED
@@ -0,0 +1,37 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # // Copyright (c) 2025 Bytedance Ltd. and/or its affiliates
2
+ # //
3
+ # // Licensed under the Apache License, Version 2.0 (the "License");
4
+ # // you may not use this file except in compliance with the License.
5
+ # // You may obtain a copy of the License at
6
+ # //
7
+ # // http://www.apache.org/licenses/LICENSE-2.0
8
+ # //
9
+ # // Unless required by applicable law or agreed to in writing, software
10
+ # // distributed under the License is distributed on an "AS IS" BASIS,
11
+ # // WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # // See the License for the specific language governing permissions and
13
+ # // limitations under the License.
14
+
15
+ """
16
+ Distributed package.
17
+ """
18
+
19
+ from .basic import (
20
+ barrier_if_distributed,
21
+ convert_to_ddp,
22
+ get_device,
23
+ get_global_rank,
24
+ get_local_rank,
25
+ get_world_size,
26
+ init_torch,
27
+ )
28
+
29
+ __all__ = [
30
+ "barrier_if_distributed",
31
+ "convert_to_ddp",
32
+ "get_device",
33
+ "get_global_rank",
34
+ "get_local_rank",
35
+ "get_world_size",
36
+ "init_torch",
37
+ ]
common/distributed/advanced.py ADDED
@@ -0,0 +1,208 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # // Copyright (c) 2025 Bytedance Ltd. and/or its affiliates
2
+ # //
3
+ # // Licensed under the Apache License, Version 2.0 (the "License");
4
+ # // you may not use this file except in compliance with the License.
5
+ # // You may obtain a copy of the License at
6
+ # //
7
+ # // http://www.apache.org/licenses/LICENSE-2.0
8
+ # //
9
+ # // Unless required by applicable law or agreed to in writing, software
10
+ # // distributed under the License is distributed on an "AS IS" BASIS,
11
+ # // WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # // See the License for the specific language governing permissions and
13
+ # // limitations under the License.
14
+
15
+ """
16
+ Advanced distributed functions for sequence parallel.
17
+ """
18
+
19
+ from typing import Optional, List
20
+ import torch
21
+ import torch.distributed as dist
22
+ from torch.distributed.device_mesh import DeviceMesh, init_device_mesh
23
+ from torch.distributed.fsdp import ShardingStrategy
24
+
25
+ from .basic import get_global_rank, get_world_size
26
+
27
+
28
+ _DATA_PARALLEL_GROUP = None
29
+ _SEQUENCE_PARALLEL_GROUP = None
30
+ _SEQUENCE_PARALLEL_CPU_GROUP = None
31
+ _MODEL_SHARD_CPU_INTER_GROUP = None
32
+ _MODEL_SHARD_CPU_INTRA_GROUP = None
33
+ _MODEL_SHARD_INTER_GROUP = None
34
+ _MODEL_SHARD_INTRA_GROUP = None
35
+ _SEQUENCE_PARALLEL_GLOBAL_RANKS = None
36
+
37
+
38
+ def get_data_parallel_group() -> Optional[dist.ProcessGroup]:
39
+ """
40
+ Get data parallel process group.
41
+ """
42
+ return _DATA_PARALLEL_GROUP
43
+
44
+
45
+ def get_sequence_parallel_group() -> Optional[dist.ProcessGroup]:
46
+ """
47
+ Get sequence parallel process group.
48
+ """
49
+ return _SEQUENCE_PARALLEL_GROUP
50
+
51
+
52
+ def get_sequence_parallel_cpu_group() -> Optional[dist.ProcessGroup]:
53
+ """
54
+ Get sequence parallel CPU process group.
55
+ """
56
+ return _SEQUENCE_PARALLEL_CPU_GROUP
57
+
58
+
59
+ def get_data_parallel_rank() -> int:
60
+ """
61
+ Get data parallel rank.
62
+ """
63
+ group = get_data_parallel_group()
64
+ return dist.get_rank(group) if group else get_global_rank()
65
+
66
+
67
+ def get_data_parallel_world_size() -> int:
68
+ """
69
+ Get data parallel world size.
70
+ """
71
+ group = get_data_parallel_group()
72
+ return dist.get_world_size(group) if group else get_world_size()
73
+
74
+
75
+ def get_sequence_parallel_rank() -> int:
76
+ """
77
+ Get sequence parallel rank.
78
+ """
79
+ group = get_sequence_parallel_group()
80
+ return dist.get_rank(group) if group else 0
81
+
82
+
83
+ def get_sequence_parallel_world_size() -> int:
84
+ """
85
+ Get sequence parallel world size.
86
+ """
87
+ group = get_sequence_parallel_group()
88
+ return dist.get_world_size(group) if group else 1
89
+
90
+
91
+ def get_model_shard_cpu_intra_group() -> Optional[dist.ProcessGroup]:
92
+ """
93
+ Get the CPU intra process group of model sharding.
94
+ """
95
+ return _MODEL_SHARD_CPU_INTRA_GROUP
96
+
97
+
98
+ def get_model_shard_cpu_inter_group() -> Optional[dist.ProcessGroup]:
99
+ """
100
+ Get the CPU inter process group of model sharding.
101
+ """
102
+ return _MODEL_SHARD_CPU_INTER_GROUP
103
+
104
+
105
+ def get_model_shard_intra_group() -> Optional[dist.ProcessGroup]:
106
+ """
107
+ Get the GPU intra process group of model sharding.
108
+ """
109
+ return _MODEL_SHARD_INTRA_GROUP
110
+
111
+
112
+ def get_model_shard_inter_group() -> Optional[dist.ProcessGroup]:
113
+ """
114
+ Get the GPU inter process group of model sharding.
115
+ """
116
+ return _MODEL_SHARD_INTER_GROUP
117
+
118
+
119
+ def init_sequence_parallel(sequence_parallel_size: int):
120
+ """
121
+ Initialize sequence parallel.
122
+ """
123
+ global _DATA_PARALLEL_GROUP
124
+ global _SEQUENCE_PARALLEL_GROUP
125
+ global _SEQUENCE_PARALLEL_CPU_GROUP
126
+ global _SEQUENCE_PARALLEL_GLOBAL_RANKS
127
+ assert dist.is_initialized()
128
+ world_size = dist.get_world_size()
129
+ rank = dist.get_rank()
130
+ data_parallel_size = world_size // sequence_parallel_size
131
+ for i in range(data_parallel_size):
132
+ start_rank = i * sequence_parallel_size
133
+ end_rank = (i + 1) * sequence_parallel_size
134
+ ranks = range(start_rank, end_rank)
135
+ group = dist.new_group(ranks)
136
+ cpu_group = dist.new_group(ranks, backend="gloo")
137
+ if rank in ranks:
138
+ _SEQUENCE_PARALLEL_GROUP = group
139
+ _SEQUENCE_PARALLEL_CPU_GROUP = cpu_group
140
+ _SEQUENCE_PARALLEL_GLOBAL_RANKS = list(ranks)
141
+
142
+
143
+ def init_model_shard_group(
144
+ *,
145
+ sharding_strategy: ShardingStrategy,
146
+ device_mesh: Optional[DeviceMesh] = None,
147
+ ):
148
+ """
149
+ Initialize process group of model sharding.
150
+ """
151
+ global _MODEL_SHARD_INTER_GROUP
152
+ global _MODEL_SHARD_INTRA_GROUP
153
+ global _MODEL_SHARD_CPU_INTER_GROUP
154
+ global _MODEL_SHARD_CPU_INTRA_GROUP
155
+ assert dist.is_initialized()
156
+ world_size = dist.get_world_size()
157
+ if device_mesh is not None:
158
+ num_shards_per_group = device_mesh.shape[1]
159
+ elif sharding_strategy == ShardingStrategy.NO_SHARD:
160
+ num_shards_per_group = 1
161
+ elif sharding_strategy in [
162
+ ShardingStrategy.HYBRID_SHARD,
163
+ ShardingStrategy._HYBRID_SHARD_ZERO2,
164
+ ]:
165
+ num_shards_per_group = torch.cuda.device_count()
166
+ else:
167
+ num_shards_per_group = world_size
168
+ num_groups = world_size // num_shards_per_group
169
+ device_mesh = (num_groups, num_shards_per_group)
170
+
171
+ gpu_mesh_2d = init_device_mesh("cuda", device_mesh, mesh_dim_names=("inter", "intra"))
172
+ cpu_mesh_2d = init_device_mesh("cpu", device_mesh, mesh_dim_names=("inter", "intra"))
173
+
174
+ _MODEL_SHARD_INTER_GROUP = gpu_mesh_2d.get_group("inter")
175
+ _MODEL_SHARD_INTRA_GROUP = gpu_mesh_2d.get_group("intra")
176
+ _MODEL_SHARD_CPU_INTER_GROUP = cpu_mesh_2d.get_group("inter")
177
+ _MODEL_SHARD_CPU_INTRA_GROUP = cpu_mesh_2d.get_group("intra")
178
+
179
+ def get_sequence_parallel_global_ranks() -> List[int]:
180
+ """
181
+ Get all global ranks of the sequence parallel process group
182
+ that the caller rank belongs to.
183
+ """
184
+ if _SEQUENCE_PARALLEL_GLOBAL_RANKS is None:
185
+ return [dist.get_rank()]
186
+ return _SEQUENCE_PARALLEL_GLOBAL_RANKS
187
+
188
+
189
+ def get_next_sequence_parallel_rank() -> int:
190
+ """
191
+ Get the next global rank of the sequence parallel process group
192
+ that the caller rank belongs to.
193
+ """
194
+ sp_global_ranks = get_sequence_parallel_global_ranks()
195
+ sp_rank = get_sequence_parallel_rank()
196
+ sp_size = get_sequence_parallel_world_size()
197
+ return sp_global_ranks[(sp_rank + 1) % sp_size]
198
+
199
+
200
+ def get_prev_sequence_parallel_rank() -> int:
201
+ """
202
+ Get the previous global rank of the sequence parallel process group
203
+ that the caller rank belongs to.
204
+ """
205
+ sp_global_ranks = get_sequence_parallel_global_ranks()
206
+ sp_rank = get_sequence_parallel_rank()
207
+ sp_size = get_sequence_parallel_world_size()
208
+ return sp_global_ranks[(sp_rank + sp_size - 1) % sp_size]
common/distributed/basic.py ADDED
@@ -0,0 +1,84 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # // Copyright (c) 2025 Bytedance Ltd. and/or its affiliates
2
+ # //
3
+ # // Licensed under the Apache License, Version 2.0 (the "License");
4
+ # // you may not use this file except in compliance with the License.
5
+ # // You may obtain a copy of the License at
6
+ # //
7
+ # // http://www.apache.org/licenses/LICENSE-2.0
8
+ # //
9
+ # // Unless required by applicable law or agreed to in writing, software
10
+ # // distributed under the License is distributed on an "AS IS" BASIS,
11
+ # // WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # // See the License for the specific language governing permissions and
13
+ # // limitations under the License.
14
+
15
+ """
16
+ Distributed basic functions.
17
+ """
18
+
19
+ import os
20
+ from datetime import timedelta
21
+ import torch
22
+ import torch.distributed as dist
23
+ from torch.nn.parallel import DistributedDataParallel
24
+
25
+
26
+ def get_global_rank() -> int:
27
+ """
28
+ Get the global rank, the global index of the GPU.
29
+ """
30
+ return int(os.environ.get("RANK", "0"))
31
+
32
+
33
+ def get_local_rank() -> int:
34
+ """
35
+ Get the local rank, the local index of the GPU.
36
+ """
37
+ return int(os.environ.get("LOCAL_RANK", "0"))
38
+
39
+
40
+ def get_world_size() -> int:
41
+ """
42
+ Get the world size, the total amount of GPUs.
43
+ """
44
+ return int(os.environ.get("WORLD_SIZE", "1"))
45
+
46
+
47
+ def get_device() -> torch.device:
48
+ """
49
+ Get current rank device.
50
+ """
51
+ return torch.device("cuda", get_local_rank())
52
+
53
+
54
+ def barrier_if_distributed(*args, **kwargs):
55
+ """
56
+ Synchronizes all processes if under distributed context.
57
+ """
58
+ if dist.is_initialized():
59
+ return dist.barrier(*args, **kwargs)
60
+
61
+
62
+ def init_torch(cudnn_benchmark=True, timeout=timedelta(seconds=600)):
63
+ """
64
+ Common PyTorch initialization configuration.
65
+ """
66
+ torch.backends.cuda.matmul.allow_tf32 = True
67
+ torch.backends.cudnn.allow_tf32 = True
68
+ torch.backends.cudnn.benchmark = cudnn_benchmark
69
+ torch.cuda.set_device(get_local_rank())
70
+ dist.init_process_group(
71
+ backend="nccl",
72
+ rank=get_global_rank(),
73
+ world_size=get_world_size(),
74
+ timeout=timeout,
75
+ )
76
+
77
+
78
+ def convert_to_ddp(module: torch.nn.Module, **kwargs) -> DistributedDataParallel:
79
+ return DistributedDataParallel(
80
+ module=module,
81
+ device_ids=[get_local_rank()],
82
+ output_device=get_local_rank(),
83
+ **kwargs,
84
+ )
common/distributed/meta_init_utils.py ADDED
@@ -0,0 +1,41 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # // Copyright (c) 2025 Bytedance Ltd. and/or its affiliates
2
+ # //
3
+ # // Licensed under the Apache License, Version 2.0 (the "License");
4
+ # // you may not use this file except in compliance with the License.
5
+ # // You may obtain a copy of the License at
6
+ # //
7
+ # // http://www.apache.org/licenses/LICENSE-2.0
8
+ # //
9
+ # // Unless required by applicable law or agreed to in writing, software
10
+ # // distributed under the License is distributed on an "AS IS" BASIS,
11
+ # // WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # // See the License for the specific language governing permissions and
13
+ # // limitations under the License.
14
+
15
+ import torch
16
+ from rotary_embedding_torch import RotaryEmbedding
17
+ from torch import nn
18
+ from torch.distributed.fsdp._common_utils import _is_fsdp_flattened
19
+
20
+ __all__ = ["meta_non_persistent_buffer_init_fn"]
21
+
22
+
23
+ def meta_non_persistent_buffer_init_fn(module: nn.Module) -> nn.Module:
24
+ """
25
+ Used for materializing `non-persistent tensor buffers` while model resuming.
26
+
27
+ Since non-persistent tensor buffers are not saved in state_dict,
28
+ when initializing model with meta device, user should materialize those buffers manually.
29
+
30
+ Currently, only `rope.dummy` is this special case.
31
+ """
32
+ with torch.no_grad():
33
+ for submodule in module.modules():
34
+ if not isinstance(submodule, RotaryEmbedding):
35
+ continue
36
+ for buffer_name, buffer in submodule.named_buffers(recurse=False):
37
+ if buffer.is_meta and "dummy" in buffer_name:
38
+ materialized_buffer = torch.zeros_like(buffer, device="cpu")
39
+ setattr(submodule, buffer_name, materialized_buffer)
40
+ assert not any(b.is_meta for n, b in module.named_buffers())
41
+ return module
common/distributed/ops.py ADDED
@@ -0,0 +1,494 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # // Copyright (c) 2025 Bytedance Ltd. and/or its affiliates
2
+ # //
3
+ # // Licensed under the Apache License, Version 2.0 (the "License");
4
+ # // you may not use this file except in compliance with the License.
5
+ # // You may obtain a copy of the License at
6
+ # //
7
+ # // http://www.apache.org/licenses/LICENSE-2.0
8
+ # //
9
+ # // Unless required by applicable law or agreed to in writing, software
10
+ # // distributed under the License is distributed on an "AS IS" BASIS,
11
+ # // WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # // See the License for the specific language governing permissions and
13
+ # // limitations under the License.
14
+
15
+ """
16
+ Distributed ops for supporting sequence parallel.
17
+ """
18
+
19
+ from collections import defaultdict
20
+ from typing import Any, Callable, Dict, List, Optional, Tuple, Union
21
+ import torch
22
+ import torch.distributed as dist
23
+ from torch import Tensor
24
+
25
+ from common.cache import Cache
26
+ from common.distributed.advanced import (
27
+ get_sequence_parallel_group,
28
+ get_sequence_parallel_rank,
29
+ get_sequence_parallel_world_size,
30
+ )
31
+
32
+ from .basic import get_device
33
+
34
+ _SEQ_DATA_BUF = defaultdict(lambda: [None, None, None])
35
+ _SEQ_DATA_META_SHAPES = defaultdict()
36
+ _SEQ_DATA_META_DTYPES = defaultdict()
37
+ _SEQ_DATA_ASYNC_COMMS = defaultdict(list)
38
+ _SYNC_BUFFER = defaultdict(dict)
39
+
40
+
41
+ def single_all_to_all(
42
+ local_input: Tensor,
43
+ scatter_dim: int,
44
+ gather_dim: int,
45
+ group: dist.ProcessGroup,
46
+ async_op: bool = False,
47
+ ):
48
+ """
49
+ A function to do all-to-all on a tensor
50
+ """
51
+ seq_world_size = dist.get_world_size(group)
52
+ prev_scatter_dim = scatter_dim
53
+ if scatter_dim != 0:
54
+ local_input = local_input.transpose(0, scatter_dim)
55
+ if gather_dim == 0:
56
+ gather_dim = scatter_dim
57
+ scatter_dim = 0
58
+
59
+ inp_shape = list(local_input.shape)
60
+ inp_shape[scatter_dim] = inp_shape[scatter_dim] // seq_world_size
61
+ input_t = local_input.reshape(
62
+ [seq_world_size, inp_shape[scatter_dim]] + inp_shape[scatter_dim + 1 :]
63
+ ).contiguous()
64
+ output = torch.empty_like(input_t)
65
+ comm = dist.all_to_all_single(output, input_t, group=group, async_op=async_op)
66
+ if async_op:
67
+ # let user's code transpose & reshape
68
+ return output, comm, prev_scatter_dim
69
+
70
+ # first dim is seq_world_size, so we can split it directly
71
+ output = torch.cat(output.split(1), dim=gather_dim + 1).squeeze(0)
72
+ if prev_scatter_dim:
73
+ output = output.transpose(0, prev_scatter_dim).contiguous()
74
+ return output
75
+
76
+
77
+ def _all_to_all(
78
+ local_input: Tensor,
79
+ scatter_dim: int,
80
+ gather_dim: int,
81
+ group: dist.ProcessGroup,
82
+ ):
83
+ seq_world_size = dist.get_world_size(group)
84
+ input_list = [
85
+ t.contiguous() for t in torch.tensor_split(local_input, seq_world_size, scatter_dim)
86
+ ]
87
+ output_list = [torch.empty_like(input_list[0]) for _ in range(seq_world_size)]
88
+ dist.all_to_all(output_list, input_list, group=group)
89
+ return torch.cat(output_list, dim=gather_dim).contiguous()
90
+
91
+
92
+ class SeqAllToAll(torch.autograd.Function):
93
+ @staticmethod
94
+ def forward(
95
+ ctx: Any,
96
+ group: dist.ProcessGroup,
97
+ local_input: Tensor,
98
+ scatter_dim: int,
99
+ gather_dim: int,
100
+ async_op: bool,
101
+ ) -> Tensor:
102
+ ctx.group = group
103
+ ctx.scatter_dim = scatter_dim
104
+ ctx.gather_dim = gather_dim
105
+ ctx.async_op = async_op
106
+ if async_op:
107
+ output, comm, prev_scatter_dim = single_all_to_all(
108
+ local_input, scatter_dim, gather_dim, group, async_op=async_op
109
+ )
110
+ ctx.prev_scatter_dim = prev_scatter_dim
111
+ return output, comm
112
+
113
+ return _all_to_all(local_input, scatter_dim, gather_dim, group)
114
+
115
+ @staticmethod
116
+ def backward(ctx: Any, *grad_output: Tensor) -> Tuple[None, Tensor, None, None]:
117
+ if ctx.async_op:
118
+ input_t = torch.cat(grad_output[0].split(1), dim=ctx.gather_dim + 1).squeeze(0)
119
+ if ctx.prev_scatter_dim:
120
+ input_t = input_t.transpose(0, ctx.prev_scatter_dim)
121
+ else:
122
+ input_t = grad_output[0]
123
+ return (
124
+ None,
125
+ _all_to_all(input_t, ctx.gather_dim, ctx.scatter_dim, ctx.group),
126
+ None,
127
+ None,
128
+ None,
129
+ )
130
+
131
+
132
+ class Slice(torch.autograd.Function):
133
+ @staticmethod
134
+ def forward(ctx: Any, group: dist.ProcessGroup, local_input: Tensor, dim: int) -> Tensor:
135
+ ctx.group = group
136
+ ctx.rank = dist.get_rank(group)
137
+ seq_world_size = dist.get_world_size(group)
138
+ ctx.seq_world_size = seq_world_size
139
+ ctx.dim = dim
140
+ dim_size = local_input.shape[dim]
141
+ return local_input.split(dim_size // seq_world_size, dim=dim)[ctx.rank].contiguous()
142
+
143
+ @staticmethod
144
+ def backward(ctx: Any, grad_output: Tensor) -> Tuple[None, Tensor, None]:
145
+ dim_size = list(grad_output.size())
146
+ split_size = dim_size[0]
147
+ dim_size[0] = dim_size[0] * ctx.seq_world_size
148
+ output = torch.empty(dim_size, dtype=grad_output.dtype, device=torch.cuda.current_device())
149
+ dist._all_gather_base(output, grad_output, group=ctx.group)
150
+ return (None, torch.cat(output.split(split_size), dim=ctx.dim), None)
151
+
152
+
153
+ class Gather(torch.autograd.Function):
154
+ @staticmethod
155
+ def forward(
156
+ ctx: Any,
157
+ group: dist.ProcessGroup,
158
+ local_input: Tensor,
159
+ dim: int,
160
+ grad_scale: Optional[bool] = False,
161
+ ) -> Tensor:
162
+ ctx.group = group
163
+ ctx.rank = dist.get_rank(group)
164
+ ctx.dim = dim
165
+ ctx.grad_scale = grad_scale
166
+ seq_world_size = dist.get_world_size(group)
167
+ ctx.seq_world_size = seq_world_size
168
+ dim_size = list(local_input.size())
169
+ split_size = dim_size[0]
170
+ ctx.part_size = dim_size[dim]
171
+ dim_size[0] = dim_size[0] * seq_world_size
172
+ output = torch.empty(dim_size, dtype=local_input.dtype, device=torch.cuda.current_device())
173
+ dist._all_gather_base(output, local_input.contiguous(), group=ctx.group)
174
+ return torch.cat(output.split(split_size), dim=dim)
175
+
176
+ @staticmethod
177
+ def backward(ctx: Any, grad_output: Tensor) -> Tuple[None, Tensor]:
178
+ if ctx.grad_scale:
179
+ grad_output = grad_output * ctx.seq_world_size
180
+ return (
181
+ None,
182
+ grad_output.split(ctx.part_size, dim=ctx.dim)[ctx.rank].contiguous(),
183
+ None,
184
+ None,
185
+ )
186
+
187
+
188
+ def gather_seq_scatter_heads_qkv(
189
+ qkv_tensor: Tensor,
190
+ *,
191
+ seq_dim: int,
192
+ qkv_shape: Optional[Tensor] = None,
193
+ cache: Cache = Cache(disable=True),
194
+ restore_shape: bool = True,
195
+ ):
196
+ """
197
+ A func to sync splited qkv tensor
198
+ qkv_tensor: the tensor we want to do alltoall with. The last dim must
199
+ be the projection_idx, which we will split into 3 part. After
200
+ spliting, the gather idx will be projecttion_idx + 1
201
+ seq_dim: gather_dim for all2all comm
202
+ restore_shape: if True, output will has the same shape length as input
203
+ """
204
+ group = get_sequence_parallel_group()
205
+ if not group:
206
+ return qkv_tensor
207
+ world = get_sequence_parallel_world_size()
208
+ orig_shape = qkv_tensor.shape
209
+ scatter_dim = qkv_tensor.dim()
210
+ bef_all2all_shape = list(orig_shape)
211
+ qkv_proj_dim = bef_all2all_shape[-1]
212
+ bef_all2all_shape = bef_all2all_shape[:-1] + [3, qkv_proj_dim // 3]
213
+ qkv_tensor = qkv_tensor.view(bef_all2all_shape)
214
+ qkv_tensor = SeqAllToAll.apply(group, qkv_tensor, scatter_dim, seq_dim, False)
215
+ if restore_shape:
216
+ out_shape = list(orig_shape)
217
+ out_shape[seq_dim] *= world
218
+ out_shape[-1] = qkv_proj_dim // world
219
+ qkv_tensor = qkv_tensor.view(out_shape)
220
+
221
+ # remove padding
222
+ if qkv_shape is not None:
223
+ unpad_dim_size = cache(
224
+ "unpad_dim_size", lambda: torch.sum(torch.prod(qkv_shape, dim=-1)).item()
225
+ )
226
+ if unpad_dim_size % world != 0:
227
+ padding_size = qkv_tensor.size(seq_dim) - unpad_dim_size
228
+ qkv_tensor = _unpad_tensor(qkv_tensor, seq_dim, padding_size)
229
+ return qkv_tensor
230
+
231
+
232
+ def slice_inputs(x: Tensor, dim: int, padding: bool = True):
233
+ """
234
+ A func to slice the input sequence in sequence parallel
235
+ """
236
+ group = get_sequence_parallel_group()
237
+ if group is None:
238
+ return x
239
+ sp_rank = get_sequence_parallel_rank()
240
+ sp_world = get_sequence_parallel_world_size()
241
+ dim_size = x.shape[dim]
242
+ unit = (dim_size + sp_world - 1) // sp_world
243
+ if padding and dim_size % sp_world:
244
+ padding_size = sp_world - (dim_size % sp_world)
245
+ x = _pad_tensor(x, dim, padding_size)
246
+ slc = [slice(None)] * len(x.shape)
247
+ slc[dim] = slice(unit * sp_rank, unit * (sp_rank + 1))
248
+ return x[slc]
249
+
250
+
251
+ def remove_seqeunce_parallel_padding(x: Tensor, dim: int, unpad_dim_size: int):
252
+ """
253
+ A func to remove the padding part of the tensor based on its original shape
254
+ """
255
+ group = get_sequence_parallel_group()
256
+ if group is None:
257
+ return x
258
+ sp_world = get_sequence_parallel_world_size()
259
+ if unpad_dim_size % sp_world == 0:
260
+ return x
261
+ padding_size = sp_world - (unpad_dim_size % sp_world)
262
+ assert (padding_size + unpad_dim_size) % sp_world == 0
263
+ return _unpad_tensor(x, dim=dim, padding_size=padding_size)
264
+
265
+
266
+ def gather_heads_scatter_seq(x: Tensor, head_dim: int, seq_dim: int) -> Tensor:
267
+ """
268
+ A func to sync attention result with alltoall in sequence parallel
269
+ """
270
+ group = get_sequence_parallel_group()
271
+ if not group:
272
+ return x
273
+ dim_size = x.size(seq_dim)
274
+ sp_world = get_sequence_parallel_world_size()
275
+ if dim_size % sp_world != 0:
276
+ padding_size = sp_world - (dim_size % sp_world)
277
+ x = _pad_tensor(x, seq_dim, padding_size)
278
+ return SeqAllToAll.apply(group, x, seq_dim, head_dim, False)
279
+
280
+
281
+ def gather_seq_scatter_heads(x: Tensor, seq_dim: int, head_dim: int) -> Tensor:
282
+ """
283
+ A func to sync embedding input with alltoall in sequence parallel
284
+ """
285
+ group = get_sequence_parallel_group()
286
+ if not group:
287
+ return x
288
+ return SeqAllToAll.apply(group, x, head_dim, seq_dim, False)
289
+
290
+
291
+ def scatter_heads(x: Tensor, dim: int) -> Tensor:
292
+ """
293
+ A func to split heads before attention in sequence parallel
294
+ """
295
+ group = get_sequence_parallel_group()
296
+ if not group:
297
+ return x
298
+ return Slice.apply(group, x, dim)
299
+
300
+
301
+ def gather_heads(x: Tensor, dim: int, grad_scale: Optional[bool] = False) -> Tensor:
302
+ """
303
+ A func to gather heads for the attention result in sequence parallel
304
+ """
305
+ group = get_sequence_parallel_group()
306
+ if not group:
307
+ return x
308
+ return Gather.apply(group, x, dim, grad_scale)
309
+
310
+
311
+ def gather_outputs(
312
+ x: Tensor,
313
+ *,
314
+ gather_dim: int,
315
+ padding_dim: Optional[int] = None,
316
+ unpad_shape: Optional[Tensor] = None,
317
+ cache: Cache = Cache(disable=True),
318
+ scale_grad=True,
319
+ ):
320
+ """
321
+ A func to gather the outputs for the model result in sequence parallel
322
+ """
323
+ group = get_sequence_parallel_group()
324
+ if not group:
325
+ return x
326
+ x = Gather.apply(group, x, gather_dim, scale_grad)
327
+ if padding_dim is not None:
328
+ unpad_dim_size = cache(
329
+ "unpad_dim_size", lambda: torch.sum(torch.prod(unpad_shape, dim=1)).item()
330
+ )
331
+ x = remove_seqeunce_parallel_padding(x, padding_dim, unpad_dim_size)
332
+ return x
333
+
334
+
335
+ def _pad_tensor(x: Tensor, dim: int, padding_size: int):
336
+ shape = list(x.shape)
337
+ shape[dim] = padding_size
338
+ pad = torch.zeros(shape, dtype=x.dtype, device=x.device)
339
+ return torch.cat([x, pad], dim=dim)
340
+
341
+
342
+ def _unpad_tensor(x: Tensor, dim: int, padding_size):
343
+ slc = [slice(None)] * len(x.shape)
344
+ slc[dim] = slice(0, -padding_size)
345
+ return x[slc]
346
+
347
+
348
+ def _broadcast_data(data, shape, dtype, src, group, async_op):
349
+ comms = []
350
+ if isinstance(data, (list, tuple)):
351
+ for i, sub_shape in enumerate(shape):
352
+ comms += _broadcast_data(data[i], sub_shape, dtype[i], src, group, async_op)
353
+ elif isinstance(data, dict):
354
+ for key, sub_data in data.items():
355
+ comms += _broadcast_data(sub_data, shape[key], dtype[key], src, group, async_op)
356
+ elif isinstance(data, Tensor):
357
+ comms.append(dist.broadcast(data, src=src, group=group, async_op=async_op))
358
+ return comms
359
+
360
+
361
+ def _traverse(data: Any, op: Callable) -> Union[None, List, Dict, Any]:
362
+ if isinstance(data, (list, tuple)):
363
+ return [_traverse(sub_data, op) for sub_data in data]
364
+ elif isinstance(data, dict):
365
+ return {key: _traverse(sub_data, op) for key, sub_data in data.items()}
366
+ elif isinstance(data, Tensor):
367
+ return op(data)
368
+ else:
369
+ return None
370
+
371
+
372
+ def _get_shapes(data):
373
+ return _traverse(data, op=lambda x: x.shape)
374
+
375
+
376
+ def _get_dtypes(data):
377
+ return _traverse(data, op=lambda x: x.dtype)
378
+
379
+
380
+ def _construct_broadcast_buffer(shapes, dtypes, device):
381
+ if isinstance(shapes, torch.Size):
382
+ return torch.empty(shapes, dtype=dtypes, device=device)
383
+
384
+ if isinstance(shapes, (list, tuple)):
385
+ buffer = []
386
+ for i, sub_shape in enumerate(shapes):
387
+ buffer.append(_construct_broadcast_buffer(sub_shape, dtypes[i], device))
388
+ elif isinstance(shapes, dict):
389
+ buffer = {}
390
+ for key, sub_shape in shapes.items():
391
+ buffer[key] = _construct_broadcast_buffer(sub_shape, dtypes[key], device)
392
+ else:
393
+ return None
394
+ return buffer
395
+
396
+
397
+ class SPDistForward:
398
+ """A forward tool to sync different result across sp group
399
+
400
+ Args:
401
+ module: a function or module to process users input
402
+ sp_step: current training step to judge which rank to broadcast its result to all
403
+ name: a distinct str to save meta and async comm
404
+ comm_shape: if different ranks have different shape, mark this arg to True
405
+ device: the device for current rank, can be empty
406
+ """
407
+
408
+ def __init__(
409
+ self,
410
+ name: str,
411
+ comm_shape: bool,
412
+ device: torch.device = None,
413
+ ):
414
+ self.name = name
415
+ self.comm_shape = comm_shape
416
+ if device:
417
+ self.device = device
418
+ else:
419
+ self.device = get_device()
420
+
421
+ def __call__(self, inputs) -> Any:
422
+ group = get_sequence_parallel_group()
423
+ if not group:
424
+ yield inputs
425
+ else:
426
+ device = self.device
427
+ sp_world = get_sequence_parallel_world_size()
428
+ sp_rank = get_sequence_parallel_rank()
429
+ for local_step in range(sp_world):
430
+ src_rank = dist.get_global_rank(group, local_step)
431
+ is_src = sp_rank == local_step
432
+ local_shapes = []
433
+ local_dtypes = []
434
+ if local_step == 0:
435
+ local_result = inputs
436
+ _SEQ_DATA_BUF[self.name][-1] = local_result
437
+ local_shapes = _get_shapes(local_result)
438
+ local_dtypes = _get_dtypes(local_result)
439
+ if self.comm_shape:
440
+ group_shapes_lists = [None] * sp_world
441
+ dist.all_gather_object(group_shapes_lists, local_shapes, group=group)
442
+ _SEQ_DATA_META_SHAPES[self.name] = group_shapes_lists
443
+ else:
444
+ _SEQ_DATA_META_SHAPES[self.name] = [local_shapes] * sp_world
445
+ _SEQ_DATA_META_DTYPES[self.name] = local_dtypes
446
+ shapes = _SEQ_DATA_META_SHAPES[self.name][local_step]
447
+ dtypes = _SEQ_DATA_META_DTYPES[self.name]
448
+ buf_id = local_step % 2
449
+ if local_step == 0:
450
+ sync_data = (
451
+ local_result
452
+ if is_src
453
+ else _construct_broadcast_buffer(shapes, dtypes, device)
454
+ )
455
+ _broadcast_data(sync_data, shapes, dtypes, src_rank, group, False)
456
+ _SEQ_DATA_BUF[self.name][buf_id] = sync_data
457
+
458
+ # wait for async comm ops
459
+ if _SEQ_DATA_ASYNC_COMMS[self.name]:
460
+ for comm in _SEQ_DATA_ASYNC_COMMS[self.name]:
461
+ comm.wait()
462
+ # before return the sync result, do async broadcast for next batch
463
+ if local_step < sp_world - 1:
464
+ next_buf_id = 1 - buf_id
465
+ shapes = _SEQ_DATA_META_SHAPES[self.name][local_step + 1]
466
+ src_rank = dist.get_global_rank(group, local_step + 1)
467
+ is_src = sp_rank == local_step + 1
468
+ next_sync_data = (
469
+ _SEQ_DATA_BUF[self.name][-1]
470
+ if is_src
471
+ else _construct_broadcast_buffer(shapes, dtypes, device)
472
+ )
473
+ _SEQ_DATA_ASYNC_COMMS[self.name] = _broadcast_data(
474
+ next_sync_data, shapes, dtypes, src_rank, group, True
475
+ )
476
+ _SEQ_DATA_BUF[self.name][next_buf_id] = next_sync_data
477
+ yield _SEQ_DATA_BUF[self.name][buf_id]
478
+
479
+
480
+ sync_inputs = SPDistForward(name="bef_fwd", comm_shape=True)
481
+
482
+
483
+ def sync_data(data, sp_idx, name="tmp"):
484
+ group = get_sequence_parallel_group()
485
+ if group is None:
486
+ return data
487
+ # if sp_idx in _SYNC_BUFFER[name]:
488
+ # return _SYNC_BUFFER[name][sp_idx]
489
+ sp_rank = get_sequence_parallel_rank()
490
+ src_rank = dist.get_global_rank(group, sp_idx)
491
+ objects = [data] if sp_rank == sp_idx else [None]
492
+ dist.broadcast_object_list(objects, src=src_rank, group=group)
493
+ # _SYNC_BUFFER[name] = {sp_idx: objects[0]}
494
+ return objects[0]
common/logger.py ADDED
@@ -0,0 +1,44 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # // Copyright (c) 2025 Bytedance Ltd. and/or its affiliates
2
+ # //
3
+ # // Licensed under the Apache License, Version 2.0 (the "License");
4
+ # // you may not use this file except in compliance with the License.
5
+ # // You may obtain a copy of the License at
6
+ # //
7
+ # // http://www.apache.org/licenses/LICENSE-2.0
8
+ # //
9
+ # // Unless required by applicable law or agreed to in writing, software
10
+ # // distributed under the License is distributed on an "AS IS" BASIS,
11
+ # // WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # // See the License for the specific language governing permissions and
13
+ # // limitations under the License.
14
+
15
+ """
16
+ Logging utility functions.
17
+ """
18
+
19
+ import logging
20
+ import sys
21
+ from typing import Optional
22
+
23
+ from common.distributed import get_global_rank, get_local_rank, get_world_size
24
+
25
+ _default_handler = logging.StreamHandler(sys.stdout)
26
+ _default_handler.setFormatter(
27
+ logging.Formatter(
28
+ "%(asctime)s "
29
+ + (f"[Rank:{get_global_rank()}]" if get_world_size() > 1 else "")
30
+ + (f"[LocalRank:{get_local_rank()}]" if get_world_size() > 1 else "")
31
+ + "[%(threadName).12s][%(name)s][%(levelname).5s] "
32
+ + "%(message)s"
33
+ )
34
+ )
35
+
36
+
37
+ def get_logger(name: Optional[str] = None) -> logging.Logger:
38
+ """
39
+ Get a logger.
40
+ """
41
+ logger = logging.getLogger(name)
42
+ logger.addHandler(_default_handler)
43
+ logger.setLevel(logging.INFO)
44
+ return logger
common/partition.py ADDED
@@ -0,0 +1,59 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # // Copyright (c) 2025 Bytedance Ltd. and/or its affiliates
2
+ # //
3
+ # // Licensed under the Apache License, Version 2.0 (the "License");
4
+ # // you may not use this file except in compliance with the License.
5
+ # // You may obtain a copy of the License at
6
+ # //
7
+ # // http://www.apache.org/licenses/LICENSE-2.0
8
+ # //
9
+ # // Unless required by applicable law or agreed to in writing, software
10
+ # // distributed under the License is distributed on an "AS IS" BASIS,
11
+ # // WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # // See the License for the specific language governing permissions and
13
+ # // limitations under the License.
14
+
15
+ """
16
+ Partition utility functions.
17
+ """
18
+
19
+ from typing import Any, List
20
+
21
+
22
+ def partition_by_size(data: List[Any], size: int) -> List[List[Any]]:
23
+ """
24
+ Partition a list by size.
25
+ When indivisible, the last group contains fewer items than the target size.
26
+
27
+ Examples:
28
+ - data: [1,2,3,4,5]
29
+ - size: 2
30
+ - return: [[1,2], [3,4], [5]]
31
+ """
32
+ assert size > 0
33
+ return [data[i : (i + size)] for i in range(0, len(data), size)]
34
+
35
+
36
+ def partition_by_groups(data: List[Any], groups: int) -> List[List[Any]]:
37
+ """
38
+ Partition a list by groups.
39
+ When indivisible, some groups may have more items than others.
40
+
41
+ Examples:
42
+ - data: [1,2,3,4,5]
43
+ - groups: 2
44
+ - return: [[1,3,5], [2,4]]
45
+ """
46
+ assert groups > 0
47
+ return [data[i::groups] for i in range(groups)]
48
+
49
+
50
+ def shift_list(data: List[Any], n: int) -> List[Any]:
51
+ """
52
+ Rotate a list by n elements.
53
+
54
+ Examples:
55
+ - data: [1,2,3,4,5]
56
+ - n: 3
57
+ - return: [4,5,1,2,3]
58
+ """
59
+ return data[(n % len(data)) :] + data[: (n % len(data))]
common/seed.py ADDED
@@ -0,0 +1,30 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # // Copyright (c) 2025 Bytedance Ltd. and/or its affiliates
2
+ # //
3
+ # // Licensed under the Apache License, Version 2.0 (the "License");
4
+ # // you may not use this file except in compliance with the License.
5
+ # // You may obtain a copy of the License at
6
+ # //
7
+ # // http://www.apache.org/licenses/LICENSE-2.0
8
+ # //
9
+ # // Unless required by applicable law or agreed to in writing, software
10
+ # // distributed under the License is distributed on an "AS IS" BASIS,
11
+ # // WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # // See the License for the specific language governing permissions and
13
+ # // limitations under the License.
14
+
15
+ import random
16
+ from typing import Optional
17
+ import numpy as np
18
+ import torch
19
+
20
+ from common.distributed import get_global_rank
21
+
22
+
23
+ def set_seed(seed: Optional[int], same_across_ranks: bool = False):
24
+ """Function that sets the seed for pseudo-random number generators."""
25
+ if seed is not None:
26
+ seed += get_global_rank() if not same_across_ranks else 0
27
+ random.seed(seed)
28
+ np.random.seed(seed)
29
+ torch.manual_seed(seed)
30
+
configs_3b/main.yaml ADDED
@@ -0,0 +1,88 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ __object__:
2
+ path: projects.video_diffusion_sr.train
3
+ name: VideoDiffusionTrainer
4
+
5
+ dit:
6
+ model:
7
+ __object__:
8
+ path: models.dit_v2.nadit
9
+ name: NaDiT
10
+ args: as_params
11
+ vid_in_channels: 33
12
+ vid_out_channels: 16
13
+ vid_dim: 2560
14
+ vid_out_norm: rms
15
+ txt_in_dim: 5120
16
+ txt_in_norm: layer
17
+ txt_dim: ${.vid_dim}
18
+ emb_dim: ${eval:'6 * ${.vid_dim}'}
19
+ heads: 20
20
+ head_dim: 128 # llm-like
21
+ expand_ratio: 4
22
+ norm: rms
23
+ norm_eps: 1.0e-05
24
+ ada: single
25
+ qk_bias: False
26
+ qk_norm: rms
27
+ patch_size: [ 1,2,2 ]
28
+ num_layers: 32 # llm-like
29
+ mm_layers: 10
30
+ mlp_type: swiglu
31
+ msa_type: None
32
+ block_type: ${eval:'${.num_layers} * ["mmdit_sr"]'} # space-full
33
+ window: ${eval:'${.num_layers} * [(4,3,3)]'} # space-full
34
+ window_method: ${eval:'${.num_layers} // 2 * ["720pwin_by_size_bysize","720pswin_by_size_bysize"]'} # space-full
35
+ rope_type: mmrope3d
36
+ rope_dim: 128
37
+ compile: False
38
+ gradient_checkpoint: True
39
+ fsdp:
40
+ sharding_strategy: _HYBRID_SHARD_ZERO2
41
+
42
+ ema:
43
+ decay: 0.9998
44
+
45
+ vae:
46
+ model:
47
+ __inherit__: models/video_vae_v3/s8_c16_t4_inflation_sd3.yaml
48
+ freeze_encoder: False
49
+ # gradient_checkpoint: True
50
+ slicing:
51
+ split_size: 4
52
+ memory_device: same
53
+ memory_limit:
54
+ conv_max_mem: 0.5
55
+ norm_max_mem: 0.5
56
+ checkpoint: ./ckpts/ema_vae.pth
57
+ scaling_factor: 0.9152
58
+ compile: False
59
+ grouping: False
60
+ dtype: bfloat16
61
+
62
+ diffusion:
63
+ schedule:
64
+ type: lerp
65
+ T: 1000.0
66
+ sampler:
67
+ type: euler
68
+ prediction_type: v_lerp
69
+ timesteps:
70
+ training:
71
+ type: logitnormal
72
+ loc: 0.0
73
+ scale: 1.0
74
+ sampling:
75
+ type: uniform_trailing
76
+ steps: 50
77
+ transform: True
78
+ loss:
79
+ type: v_lerp
80
+ cfg:
81
+ scale: 7.5
82
+ rescale: 0
83
+
84
+ condition:
85
+ i2v: 0.0
86
+ v2v: 0.0
87
+ sr: 1.0
88
+ noise_scale: 0.25
data/image/transforms/area_resize.py ADDED
@@ -0,0 +1,135 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # // Copyright (c) 2025 Bytedance Ltd. and/or its affiliates
2
+ # //
3
+ # // Licensed under the Apache License, Version 2.0 (the "License");
4
+ # // you may not use this file except in compliance with the License.
5
+ # // You may obtain a copy of the License at
6
+ # //
7
+ # // http://www.apache.org/licenses/LICENSE-2.0
8
+ # //
9
+ # // Unless required by applicable law or agreed to in writing, software
10
+ # // distributed under the License is distributed on an "AS IS" BASIS,
11
+ # // WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # // See the License for the specific language governing permissions and
13
+ # // limitations under the License.
14
+
15
+ import math
16
+ import random
17
+ from typing import Union
18
+ import torch
19
+ from PIL import Image
20
+ from torchvision.transforms import functional as TVF
21
+ from torchvision.transforms.functional import InterpolationMode
22
+
23
+
24
+ class AreaResize:
25
+ def __init__(
26
+ self,
27
+ max_area: float,
28
+ downsample_only: bool = False,
29
+ interpolation: InterpolationMode = InterpolationMode.BICUBIC,
30
+ ):
31
+ self.max_area = max_area
32
+ self.downsample_only = downsample_only
33
+ self.interpolation = interpolation
34
+
35
+ def __call__(self, image: Union[torch.Tensor, Image.Image]):
36
+
37
+ if isinstance(image, torch.Tensor):
38
+ height, width = image.shape[-2:]
39
+ elif isinstance(image, Image.Image):
40
+ width, height = image.size
41
+ else:
42
+ raise NotImplementedError
43
+
44
+ scale = math.sqrt(self.max_area / (height * width))
45
+
46
+ # keep original height and width for small pictures.
47
+ scale = 1 if scale >= 1 and self.downsample_only else scale
48
+
49
+ resized_height, resized_width = round(height * scale), round(width * scale)
50
+
51
+ return TVF.resize(
52
+ image,
53
+ size=(resized_height, resized_width),
54
+ interpolation=self.interpolation,
55
+ )
56
+
57
+
58
+ class AreaRandomCrop:
59
+ def __init__(
60
+ self,
61
+ max_area: float,
62
+ ):
63
+ self.max_area = max_area
64
+
65
+ def get_params(self, input_size, output_size):
66
+ """Get parameters for ``crop`` for a random crop.
67
+
68
+ Args:
69
+ img (PIL Image): Image to be cropped.
70
+ output_size (tuple): Expected output size of the crop.
71
+
72
+ Returns:
73
+ tuple: params (i, j, h, w) to be passed to ``crop`` for random crop.
74
+ """
75
+ # w, h = _get_image_size(img)
76
+ h, w = input_size
77
+ th, tw = output_size
78
+ if w <= tw and h <= th:
79
+ return 0, 0, h, w
80
+
81
+ i = random.randint(0, h - th)
82
+ j = random.randint(0, w - tw)
83
+ return i, j, th, tw
84
+
85
+ def __call__(self, image: Union[torch.Tensor, Image.Image]):
86
+ if isinstance(image, torch.Tensor):
87
+ height, width = image.shape[-2:]
88
+ elif isinstance(image, Image.Image):
89
+ width, height = image.size
90
+ else:
91
+ raise NotImplementedError
92
+
93
+ resized_height = math.sqrt(self.max_area / (width / height))
94
+ resized_width = (width / height) * resized_height
95
+
96
+ # print('>>>>>>>>>>>>>>>>>>>>>')
97
+ # print((height, width))
98
+ # print( (resized_height, resized_width))
99
+
100
+ resized_height, resized_width = round(resized_height), round(resized_width)
101
+ i, j, h, w = self.get_params((height, width), (resized_height, resized_width))
102
+ image = TVF.crop(image, i, j, h, w)
103
+ return image
104
+
105
+ class ScaleResize:
106
+ def __init__(
107
+ self,
108
+ scale: float,
109
+ ):
110
+ self.scale = scale
111
+
112
+ def __call__(self, image: Union[torch.Tensor, Image.Image]):
113
+ if isinstance(image, torch.Tensor):
114
+ height, width = image.shape[-2:]
115
+ interpolation_mode = InterpolationMode.BILINEAR
116
+ antialias = True if image.ndim == 4 else "warn"
117
+ elif isinstance(image, Image.Image):
118
+ width, height = image.size
119
+ interpolation_mode = InterpolationMode.LANCZOS
120
+ antialias = "warn"
121
+ else:
122
+ raise NotImplementedError
123
+
124
+ scale = self.scale
125
+
126
+ # keep original height and width for small pictures
127
+
128
+ resized_height, resized_width = round(height * scale), round(width * scale)
129
+ image = TVF.resize(
130
+ image,
131
+ size=(resized_height, resized_width),
132
+ interpolation=interpolation_mode,
133
+ antialias=antialias,
134
+ )
135
+ return image
data/image/transforms/divisible_crop.py ADDED
@@ -0,0 +1,40 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # // Copyright (c) 2025 Bytedance Ltd. and/or its affiliates
2
+ # //
3
+ # // Licensed under the Apache License, Version 2.0 (the "License");
4
+ # // you may not use this file except in compliance with the License.
5
+ # // You may obtain a copy of the License at
6
+ # //
7
+ # // http://www.apache.org/licenses/LICENSE-2.0
8
+ # //
9
+ # // Unless required by applicable law or agreed to in writing, software
10
+ # // distributed under the License is distributed on an "AS IS" BASIS,
11
+ # // WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # // See the License for the specific language governing permissions and
13
+ # // limitations under the License.
14
+
15
+ from typing import Union
16
+ import torch
17
+ from PIL import Image
18
+ from torchvision.transforms import functional as TVF
19
+
20
+
21
+ class DivisibleCrop:
22
+ def __init__(self, factor):
23
+ if not isinstance(factor, tuple):
24
+ factor = (factor, factor)
25
+
26
+ self.height_factor, self.width_factor = factor[0], factor[1]
27
+
28
+ def __call__(self, image: Union[torch.Tensor, Image.Image]):
29
+ if isinstance(image, torch.Tensor):
30
+ height, width = image.shape[-2:]
31
+ elif isinstance(image, Image.Image):
32
+ width, height = image.size
33
+ else:
34
+ raise NotImplementedError
35
+
36
+ cropped_height = height - (height % self.height_factor)
37
+ cropped_width = width - (width % self.width_factor)
38
+
39
+ image = TVF.center_crop(img=image, output_size=(cropped_height, cropped_width))
40
+ return image
data/image/transforms/na_resize.py ADDED
@@ -0,0 +1,50 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # // Copyright (c) 2025 Bytedance Ltd. and/or its affiliates
2
+ # //
3
+ # // Licensed under the Apache License, Version 2.0 (the "License");
4
+ # // you may not use this file except in compliance with the License.
5
+ # // You may obtain a copy of the License at
6
+ # //
7
+ # // http://www.apache.org/licenses/LICENSE-2.0
8
+ # //
9
+ # // Unless required by applicable law or agreed to in writing, software
10
+ # // distributed under the License is distributed on an "AS IS" BASIS,
11
+ # // WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # // See the License for the specific language governing permissions and
13
+ # // limitations under the License.
14
+
15
+ from typing import Literal
16
+ from torchvision.transforms import CenterCrop, Compose, InterpolationMode, Resize
17
+
18
+ from .area_resize import AreaResize
19
+ from .side_resize import SideResize
20
+
21
+
22
+ def NaResize(
23
+ resolution: int,
24
+ mode: Literal["area", "side"],
25
+ downsample_only: bool,
26
+ interpolation: InterpolationMode = InterpolationMode.BICUBIC,
27
+ ):
28
+ if mode == "area":
29
+ return AreaResize(
30
+ max_area=resolution**2,
31
+ downsample_only=downsample_only,
32
+ interpolation=interpolation,
33
+ )
34
+ if mode == "side":
35
+ return SideResize(
36
+ size=resolution,
37
+ downsample_only=downsample_only,
38
+ interpolation=interpolation,
39
+ )
40
+ if mode == "square":
41
+ return Compose(
42
+ [
43
+ Resize(
44
+ size=resolution,
45
+ interpolation=interpolation,
46
+ ),
47
+ CenterCrop(resolution),
48
+ ]
49
+ )
50
+ raise ValueError(f"Unknown resize mode: {mode}")
data/image/transforms/side_resize.py ADDED
@@ -0,0 +1,54 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # // Copyright (c) 2025 Bytedance Ltd. and/or its affiliates
2
+ # //
3
+ # // Licensed under the Apache License, Version 2.0 (the "License");
4
+ # // you may not use this file except in compliance with the License.
5
+ # // You may obtain a copy of the License at
6
+ # //
7
+ # // http://www.apache.org/licenses/LICENSE-2.0
8
+ # //
9
+ # // Unless required by applicable law or agreed to in writing, software
10
+ # // distributed under the License is distributed on an "AS IS" BASIS,
11
+ # // WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # // See the License for the specific language governing permissions and
13
+ # // limitations under the License.
14
+
15
+ from typing import Union
16
+ import torch
17
+ from PIL import Image
18
+ from torchvision.transforms import InterpolationMode
19
+ from torchvision.transforms import functional as TVF
20
+
21
+
22
+ class SideResize:
23
+ def __init__(
24
+ self,
25
+ size: int,
26
+ downsample_only: bool = False,
27
+ interpolation: InterpolationMode = InterpolationMode.BICUBIC,
28
+ ):
29
+ self.size = size
30
+ self.downsample_only = downsample_only
31
+ self.interpolation = interpolation
32
+
33
+ def __call__(self, image: Union[torch.Tensor, Image.Image]):
34
+ """
35
+ Args:
36
+ image (PIL Image or Tensor): Image to be scaled.
37
+
38
+ Returns:
39
+ PIL Image or Tensor: Rescaled image.
40
+ """
41
+ if isinstance(image, torch.Tensor):
42
+ height, width = image.shape[-2:]
43
+ elif isinstance(image, Image.Image):
44
+ width, height = image.size
45
+ else:
46
+ raise NotImplementedError
47
+
48
+ if self.downsample_only and min(width, height) < self.size:
49
+ # keep original height and width for small pictures.
50
+ size = min(width, height)
51
+ else:
52
+ size = self.size
53
+
54
+ return TVF.resize(image, size, self.interpolation)
data/video/transforms/rearrange.py ADDED
@@ -0,0 +1,24 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # // Copyright (c) 2025 Bytedance Ltd. and/or its affiliates
2
+ # //
3
+ # // Licensed under the Apache License, Version 2.0 (the "License");
4
+ # // you may not use this file except in compliance with the License.
5
+ # // You may obtain a copy of the License at
6
+ # //
7
+ # // http://www.apache.org/licenses/LICENSE-2.0
8
+ # //
9
+ # // Unless required by applicable law or agreed to in writing, software
10
+ # // distributed under the License is distributed on an "AS IS" BASIS,
11
+ # // WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # // See the License for the specific language governing permissions and
13
+ # // limitations under the License.
14
+
15
+ from einops import rearrange
16
+
17
+
18
+ class Rearrange:
19
+ def __init__(self, pattern: str, **kwargs):
20
+ self.pattern = pattern
21
+ self.kwargs = kwargs
22
+
23
+ def __call__(self, x):
24
+ return rearrange(x, self.pattern, **self.kwargs)
models/dit_v2/attention.py ADDED
@@ -0,0 +1,86 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # // Copyright (c) 2025 Bytedance Ltd. and/or its affiliates
2
+ # //
3
+ # // Licensed under the Apache License, Version 2.0 (the "License");
4
+ # // you may not use this file except in compliance with the License.
5
+ # // You may obtain a copy of the License at
6
+ # //
7
+ # // http://www.apache.org/licenses/LICENSE-2.0
8
+ # //
9
+ # // Unless required by applicable law or agreed to in writing, software
10
+ # // distributed under the License is distributed on an "AS IS" BASIS,
11
+ # // WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # // See the License for the specific language governing permissions and
13
+ # // limitations under the License.
14
+
15
+ import torch
16
+ import torch.nn.functional as F
17
+
18
+ from torch import nn
19
+
20
+ # CHANGED FROM UPSTREAM (Upsampler): flash-attn is imported lazily instead of at
21
+ # module scope, and falls back to PyTorch SDPA when it is unavailable.
22
+ #
23
+ # Upstream ships a prebuilt `apex-0.1-cp310-...whl` and expects a flash-attn
24
+ # built against torch 2.4. ZeroGPU runs Python 3.12 on torch 2.8+, where neither
25
+ # installs, so the module-scope import took the whole Space down at startup.
26
+ try:
27
+ from flash_attn import flash_attn_varlen_func
28
+ except ImportError: # pragma: no cover - depends on the runtime image
29
+ flash_attn_varlen_func = None
30
+
31
+
32
+ def _sdpa_varlen(q, k, v, cu_seqlens_q, cu_seqlens_k, max_seqlen_q, max_seqlen_k, **_):
33
+ """SDPA equivalent of `flash_attn_varlen_func` for this model's usage.
34
+
35
+ Only the shape this DiT actually calls is supported, and every one of those
36
+ constraints is asserted rather than assumed: packed `(total_len, heads,
37
+ head_dim)` self-attention, identical q/k segmentation, no causal mask, no
38
+ dropout, and the default softmax scale. Sequences are unpacked from
39
+ `cu_seqlens`, attended independently, and repacked — with a batch of one
40
+ (every request this Space serves) that is a single unmasked SDPA call, so
41
+ the fallback is not a slow path so much as a direct one.
42
+ """
43
+ assert torch.equal(cu_seqlens_q, cu_seqlens_k), (
44
+ "the SDPA fallback assumes self-attention with one segmentation"
45
+ )
46
+ lengths = (cu_seqlens_q[1:] - cu_seqlens_q[:-1]).tolist()
47
+ out = []
48
+ start = 0
49
+ for length in lengths:
50
+ end = start + length
51
+ # (l h d) -> (1 h l d), attend, then back to (l h d).
52
+ segment = [
53
+ x[start:end].transpose(0, 1).unsqueeze(0) for x in (q, k, v)
54
+ ]
55
+ attended = F.scaled_dot_product_attention(*segment)
56
+ out.append(attended.squeeze(0).transpose(0, 1))
57
+ start = end
58
+ return torch.cat(out, dim=0)
59
+
60
+ class TorchAttention(nn.Module):
61
+ def tflops(self, args, kwargs, output) -> float:
62
+ assert len(args) == 0 or len(args) > 2, "query, key should both provided by args / kwargs"
63
+ q = kwargs.get("query") or args[0]
64
+ k = kwargs.get("key") or args[1]
65
+ b, h, sq, d = q.shape
66
+ b, h, sk, d = k.shape
67
+ return b * h * (4 * d * (sq / 1e6) * (sk / 1e6))
68
+
69
+ def forward(self, *args, **kwargs):
70
+ return F.scaled_dot_product_attention(*args, **kwargs)
71
+
72
+
73
+ class FlashAttentionVarlen(nn.Module):
74
+ def tflops(self, args, kwargs, output) -> float:
75
+ cu_seqlens_q = kwargs["cu_seqlens_q"]
76
+ cu_seqlens_k = kwargs["cu_seqlens_k"]
77
+ _, h, d = output.shape
78
+ seqlens_q = (cu_seqlens_q[1:] - cu_seqlens_q[:-1]) / 1e6
79
+ seqlens_k = (cu_seqlens_k[1:] - cu_seqlens_k[:-1]) / 1e6
80
+ return h * (4 * d * (seqlens_q * seqlens_k).sum())
81
+
82
+ def forward(self, *args, **kwargs):
83
+ if flash_attn_varlen_func is None:
84
+ return _sdpa_varlen(*args, **kwargs)
85
+ kwargs["deterministic"] = torch.are_deterministic_algorithms_enabled()
86
+ return flash_attn_varlen_func(*args, **kwargs)
models/dit_v2/embedding.py ADDED
@@ -0,0 +1,62 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # // Copyright (c) 2025 Bytedance Ltd. and/or its affiliates
2
+ # //
3
+ # // Licensed under the Apache License, Version 2.0 (the "License");
4
+ # // you may not use this file except in compliance with the License.
5
+ # // You may obtain a copy of the License at
6
+ # //
7
+ # // http://www.apache.org/licenses/LICENSE-2.0
8
+ # //
9
+ # // Unless required by applicable law or agreed to in writing, software
10
+ # // distributed under the License is distributed on an "AS IS" BASIS,
11
+ # // WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # // See the License for the specific language governing permissions and
13
+ # // limitations under the License.
14
+
15
+ from typing import Optional, Union
16
+ import torch
17
+ from diffusers.models.embeddings import get_timestep_embedding
18
+ from torch import nn
19
+
20
+
21
+ def emb_add(emb1: torch.Tensor, emb2: Optional[torch.Tensor]):
22
+ return emb1 if emb2 is None else emb1 + emb2
23
+
24
+
25
+ class TimeEmbedding(nn.Module):
26
+ def __init__(
27
+ self,
28
+ sinusoidal_dim: int,
29
+ hidden_dim: int,
30
+ output_dim: int,
31
+ ):
32
+ super().__init__()
33
+ self.sinusoidal_dim = sinusoidal_dim
34
+ self.proj_in = nn.Linear(sinusoidal_dim, hidden_dim)
35
+ self.proj_hid = nn.Linear(hidden_dim, hidden_dim)
36
+ self.proj_out = nn.Linear(hidden_dim, output_dim)
37
+ self.act = nn.SiLU()
38
+
39
+ def forward(
40
+ self,
41
+ timestep: Union[int, float, torch.IntTensor, torch.FloatTensor],
42
+ device: torch.device,
43
+ dtype: torch.dtype,
44
+ ) -> torch.FloatTensor:
45
+ if not torch.is_tensor(timestep):
46
+ timestep = torch.tensor([timestep], device=device, dtype=dtype)
47
+ if timestep.ndim == 0:
48
+ timestep = timestep[None]
49
+
50
+ emb = get_timestep_embedding(
51
+ timesteps=timestep,
52
+ embedding_dim=self.sinusoidal_dim,
53
+ flip_sin_to_cos=False,
54
+ downscale_freq_shift=0,
55
+ )
56
+ emb = emb.to(dtype)
57
+ emb = self.proj_in(emb)
58
+ emb = self.act(emb)
59
+ emb = self.proj_hid(emb)
60
+ emb = self.act(emb)
61
+ emb = self.proj_out(emb)
62
+ return emb
models/dit_v2/mlp.py ADDED
@@ -0,0 +1,62 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # // Copyright (c) 2025 Bytedance Ltd. and/or its affiliates
2
+ # //
3
+ # // Licensed under the Apache License, Version 2.0 (the "License");
4
+ # // you may not use this file except in compliance with the License.
5
+ # // You may obtain a copy of the License at
6
+ # //
7
+ # // http://www.apache.org/licenses/LICENSE-2.0
8
+ # //
9
+ # // Unless required by applicable law or agreed to in writing, software
10
+ # // distributed under the License is distributed on an "AS IS" BASIS,
11
+ # // WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # // See the License for the specific language governing permissions and
13
+ # // limitations under the License.
14
+
15
+ from typing import Optional
16
+ import torch
17
+ import torch.nn.functional as F
18
+ from torch import nn
19
+
20
+
21
+ def get_mlp(mlp_type: Optional[str] = "normal"):
22
+ if mlp_type == "normal":
23
+ return MLP
24
+ elif mlp_type == "swiglu":
25
+ return SwiGLUMLP
26
+
27
+
28
+ class MLP(nn.Module):
29
+ def __init__(
30
+ self,
31
+ dim: int,
32
+ expand_ratio: int,
33
+ ):
34
+ super().__init__()
35
+ self.proj_in = nn.Linear(dim, dim * expand_ratio)
36
+ self.act = nn.GELU("tanh")
37
+ self.proj_out = nn.Linear(dim * expand_ratio, dim)
38
+
39
+ def forward(self, x: torch.FloatTensor) -> torch.FloatTensor:
40
+ x = self.proj_in(x)
41
+ x = self.act(x)
42
+ x = self.proj_out(x)
43
+ return x
44
+
45
+
46
+ class SwiGLUMLP(nn.Module):
47
+ def __init__(
48
+ self,
49
+ dim: int,
50
+ expand_ratio: int,
51
+ multiple_of: int = 256,
52
+ ):
53
+ super().__init__()
54
+ hidden_dim = int(2 * dim * expand_ratio / 3)
55
+ hidden_dim = multiple_of * ((hidden_dim + multiple_of - 1) // multiple_of)
56
+ self.proj_in_gate = nn.Linear(dim, hidden_dim, bias=False)
57
+ self.proj_out = nn.Linear(hidden_dim, dim, bias=False)
58
+ self.proj_in = nn.Linear(dim, hidden_dim, bias=False)
59
+
60
+ def forward(self, x: torch.FloatTensor) -> torch.FloatTensor:
61
+ x = self.proj_out(F.silu(self.proj_in_gate(x)) * self.proj_in(x))
62
+ return x
models/dit_v2/mm.py ADDED
@@ -0,0 +1,74 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # // Copyright (c) 2025 Bytedance Ltd. and/or its affiliates
2
+ # //
3
+ # // Licensed under the Apache License, Version 2.0 (the "License");
4
+ # // you may not use this file except in compliance with the License.
5
+ # // You may obtain a copy of the License at
6
+ # //
7
+ # // http://www.apache.org/licenses/LICENSE-2.0
8
+ # //
9
+ # // Unless required by applicable law or agreed to in writing, software
10
+ # // distributed under the License is distributed on an "AS IS" BASIS,
11
+ # // WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # // See the License for the specific language governing permissions and
13
+ # // limitations under the License.
14
+
15
+ from dataclasses import dataclass
16
+ from typing import Any, Callable, Dict, List, Tuple
17
+ import torch
18
+ from torch import nn
19
+
20
+
21
+ @dataclass
22
+ class MMArg:
23
+ vid: Any
24
+ txt: Any
25
+
26
+
27
+ def get_args(key: str, args: List[Any]) -> List[Any]:
28
+ return [getattr(v, key) if isinstance(v, MMArg) else v for v in args]
29
+
30
+
31
+ def get_kwargs(key: str, kwargs: Dict[str, Any]) -> Dict[str, Any]:
32
+ return {k: getattr(v, key) if isinstance(v, MMArg) else v for k, v in kwargs.items()}
33
+
34
+
35
+ class MMModule(nn.Module):
36
+ def __init__(
37
+ self,
38
+ module: Callable[..., nn.Module],
39
+ *args,
40
+ shared_weights: bool = False,
41
+ vid_only: bool = False,
42
+ **kwargs,
43
+ ):
44
+ super().__init__()
45
+ self.shared_weights = shared_weights
46
+ self.vid_only = vid_only
47
+ if self.shared_weights:
48
+ assert get_args("vid", args) == get_args("txt", args)
49
+ assert get_kwargs("vid", kwargs) == get_kwargs("txt", kwargs)
50
+ self.all = module(*get_args("vid", args), **get_kwargs("vid", kwargs))
51
+ else:
52
+ self.vid = module(*get_args("vid", args), **get_kwargs("vid", kwargs))
53
+ self.txt = (
54
+ module(*get_args("txt", args), **get_kwargs("txt", kwargs))
55
+ if not vid_only
56
+ else None
57
+ )
58
+
59
+ def forward(
60
+ self,
61
+ vid: torch.FloatTensor,
62
+ txt: torch.FloatTensor,
63
+ *args,
64
+ **kwargs,
65
+ ) -> Tuple[
66
+ torch.FloatTensor,
67
+ torch.FloatTensor,
68
+ ]:
69
+ vid_module = self.vid if not self.shared_weights else self.all
70
+ vid = vid_module(vid, *get_args("vid", args), **get_kwargs("vid", kwargs))
71
+ if not self.vid_only:
72
+ txt_module = self.txt if not self.shared_weights else self.all
73
+ txt = txt_module(txt, *get_args("txt", args), **get_kwargs("txt", kwargs))
74
+ return vid, txt
models/dit_v2/modulation.py ADDED
@@ -0,0 +1,102 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # // Copyright (c) 2025 Bytedance Ltd. and/or its affiliates
2
+ # //
3
+ # // Licensed under the Apache License, Version 2.0 (the "License");
4
+ # // you may not use this file except in compliance with the License.
5
+ # // You may obtain a copy of the License at
6
+ # //
7
+ # // http://www.apache.org/licenses/LICENSE-2.0
8
+ # //
9
+ # // Unless required by applicable law or agreed to in writing, software
10
+ # // distributed under the License is distributed on an "AS IS" BASIS,
11
+ # // WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # // See the License for the specific language governing permissions and
13
+ # // limitations under the License.
14
+
15
+ from typing import Callable, List, Optional
16
+ import torch
17
+ from einops import rearrange
18
+ from torch import nn
19
+
20
+ from common.cache import Cache
21
+ from common.distributed.ops import slice_inputs
22
+
23
+ # (dim: int, emb_dim: int)
24
+ ada_layer_type = Callable[[int, int], nn.Module]
25
+
26
+
27
+ def get_ada_layer(ada_layer: str) -> ada_layer_type:
28
+ if ada_layer == "single":
29
+ return AdaSingle
30
+ raise NotImplementedError(f"{ada_layer} is not supported")
31
+
32
+
33
+ def expand_dims(x: torch.Tensor, dim: int, ndim: int):
34
+ """
35
+ Expand tensor "x" to "ndim" by adding empty dims at "dim".
36
+ Example: x is (b d), target ndim is 5, add dim at 1, return (b 1 1 1 d).
37
+ """
38
+ shape = x.shape
39
+ shape = shape[:dim] + (1,) * (ndim - len(shape)) + shape[dim:]
40
+ return x.reshape(shape)
41
+
42
+
43
+ class AdaSingle(nn.Module):
44
+ def __init__(
45
+ self,
46
+ dim: int,
47
+ emb_dim: int,
48
+ layers: List[str],
49
+ modes: List[str] = ["in", "out"],
50
+ ):
51
+ assert emb_dim == 6 * dim, "AdaSingle requires emb_dim == 6 * dim"
52
+ super().__init__()
53
+ self.dim = dim
54
+ self.emb_dim = emb_dim
55
+ self.layers = layers
56
+ for l in layers:
57
+ if "in" in modes:
58
+ self.register_parameter(f"{l}_shift", nn.Parameter(torch.randn(dim) / dim**0.5))
59
+ self.register_parameter(
60
+ f"{l}_scale", nn.Parameter(torch.randn(dim) / dim**0.5 + 1)
61
+ )
62
+ if "out" in modes:
63
+ self.register_parameter(f"{l}_gate", nn.Parameter(torch.randn(dim) / dim**0.5))
64
+
65
+ def forward(
66
+ self,
67
+ hid: torch.FloatTensor, # b ... c
68
+ emb: torch.FloatTensor, # b d
69
+ layer: str,
70
+ mode: str,
71
+ cache: Cache = Cache(disable=True),
72
+ branch_tag: str = "",
73
+ hid_len: Optional[torch.LongTensor] = None, # b
74
+ ) -> torch.FloatTensor:
75
+ idx = self.layers.index(layer)
76
+ emb = rearrange(emb, "b (d l g) -> b d l g", l=len(self.layers), g=3)[..., idx, :]
77
+ emb = expand_dims(emb, 1, hid.ndim + 1)
78
+
79
+ if hid_len is not None:
80
+ emb = cache(
81
+ f"emb_repeat_{idx}_{branch_tag}",
82
+ lambda: slice_inputs(
83
+ torch.cat([e.repeat(l, *([1] * e.ndim)) for e, l in zip(emb, hid_len)]),
84
+ dim=0,
85
+ ),
86
+ )
87
+
88
+ shiftA, scaleA, gateA = emb.unbind(-1)
89
+ shiftB, scaleB, gateB = (
90
+ getattr(self, f"{layer}_shift", None),
91
+ getattr(self, f"{layer}_scale", None),
92
+ getattr(self, f"{layer}_gate", None),
93
+ )
94
+
95
+ if mode == "in":
96
+ return hid.mul_(scaleA + scaleB).add_(shiftA + shiftB)
97
+ if mode == "out":
98
+ return hid.mul_(gateA + gateB)
99
+ raise NotImplementedError
100
+
101
+ def extra_repr(self) -> str:
102
+ return f"dim={self.dim}, emb_dim={self.emb_dim}, layers={self.layers}"
models/dit_v2/na.py ADDED
@@ -0,0 +1,241 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # // Copyright (c) 2025 Bytedance Ltd. and/or its affiliates
2
+ # //
3
+ # // Licensed under the Apache License, Version 2.0 (the "License");
4
+ # // you may not use this file except in compliance with the License.
5
+ # // You may obtain a copy of the License at
6
+ # //
7
+ # // http://www.apache.org/licenses/LICENSE-2.0
8
+ # //
9
+ # // Unless required by applicable law or agreed to in writing, software
10
+ # // distributed under the License is distributed on an "AS IS" BASIS,
11
+ # // WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # // See the License for the specific language governing permissions and
13
+ # // limitations under the License.
14
+
15
+ from itertools import chain
16
+ from typing import Callable, Dict, List, Tuple
17
+ import einops
18
+ import torch
19
+
20
+
21
+ def flatten(
22
+ hid: List[torch.FloatTensor], # List of (*** c)
23
+ ) -> Tuple[
24
+ torch.FloatTensor, # (L c)
25
+ torch.LongTensor, # (b n)
26
+ ]:
27
+ assert len(hid) > 0
28
+ shape = torch.stack([torch.tensor(x.shape[:-1], device=hid[0].device) for x in hid])
29
+ hid = torch.cat([x.flatten(0, -2) for x in hid])
30
+ return hid, shape
31
+
32
+
33
+ def unflatten(
34
+ hid: torch.FloatTensor, # (L c) or (L ... c)
35
+ hid_shape: torch.LongTensor, # (b n)
36
+ ) -> List[torch.Tensor]: # List of (*** c) or (*** ... c)
37
+ hid_len = hid_shape.prod(-1)
38
+ hid = hid.split(hid_len.tolist())
39
+ hid = [x.unflatten(0, s.tolist()) for x, s in zip(hid, hid_shape)]
40
+ return hid
41
+
42
+
43
+ def concat(
44
+ vid: torch.FloatTensor, # (VL ... c)
45
+ txt: torch.FloatTensor, # (TL ... c)
46
+ vid_len: torch.LongTensor, # (b)
47
+ txt_len: torch.LongTensor, # (b)
48
+ ) -> torch.FloatTensor: # (L ... c)
49
+ vid = torch.split(vid, vid_len.tolist())
50
+ txt = torch.split(txt, txt_len.tolist())
51
+ return torch.cat(list(chain(*zip(vid, txt))))
52
+
53
+
54
+ def concat_idx(
55
+ vid_len: torch.LongTensor, # (b)
56
+ txt_len: torch.LongTensor, # (b)
57
+ ) -> Tuple[
58
+ Callable,
59
+ Callable,
60
+ ]:
61
+ device = vid_len.device
62
+ vid_idx = torch.arange(vid_len.sum(), device=device)
63
+ txt_idx = torch.arange(len(vid_idx), len(vid_idx) + txt_len.sum(), device=device)
64
+ tgt_idx = concat(vid_idx, txt_idx, vid_len, txt_len)
65
+ src_idx = torch.argsort(tgt_idx)
66
+ return (
67
+ lambda vid, txt: torch.index_select(torch.cat([vid, txt]), 0, tgt_idx),
68
+ lambda all: torch.index_select(all, 0, src_idx).split([len(vid_idx), len(txt_idx)]),
69
+ )
70
+
71
+
72
+ def unconcat(
73
+ all: torch.FloatTensor, # (L ... c)
74
+ vid_len: torch.LongTensor, # (b)
75
+ txt_len: torch.LongTensor, # (b)
76
+ ) -> Tuple[
77
+ torch.FloatTensor, # (VL ... c)
78
+ torch.FloatTensor, # (TL ... c)
79
+ ]:
80
+ interleave_len = list(chain(*zip(vid_len.tolist(), txt_len.tolist())))
81
+ all = all.split(interleave_len)
82
+ vid = torch.cat(all[0::2])
83
+ txt = torch.cat(all[1::2])
84
+ return vid, txt
85
+
86
+
87
+ def repeat_concat(
88
+ vid: torch.FloatTensor, # (VL ... c)
89
+ txt: torch.FloatTensor, # (TL ... c)
90
+ vid_len: torch.LongTensor, # (n*b)
91
+ txt_len: torch.LongTensor, # (b)
92
+ txt_repeat: List, # (n)
93
+ ) -> torch.FloatTensor: # (L ... c)
94
+ vid = torch.split(vid, vid_len.tolist())
95
+ txt = torch.split(txt, txt_len.tolist())
96
+ txt = [[x] * n for x, n in zip(txt, txt_repeat)]
97
+ txt = list(chain(*txt))
98
+ return torch.cat(list(chain(*zip(vid, txt))))
99
+
100
+
101
+ def repeat_concat_idx(
102
+ vid_len: torch.LongTensor, # (n*b)
103
+ txt_len: torch.LongTensor, # (b)
104
+ txt_repeat: torch.LongTensor, # (n)
105
+ ) -> Tuple[
106
+ Callable,
107
+ Callable,
108
+ ]:
109
+ device = vid_len.device
110
+ vid_idx = torch.arange(vid_len.sum(), device=device)
111
+ txt_idx = torch.arange(len(vid_idx), len(vid_idx) + txt_len.sum(), device=device)
112
+ txt_repeat_list = txt_repeat.tolist()
113
+ tgt_idx = repeat_concat(vid_idx, txt_idx, vid_len, txt_len, txt_repeat)
114
+ src_idx = torch.argsort(tgt_idx)
115
+ txt_idx_len = len(tgt_idx) - len(vid_idx)
116
+ repeat_txt_len = (txt_len * txt_repeat).tolist()
117
+
118
+ def unconcat_coalesce(all):
119
+ """
120
+ Un-concat vid & txt, and coalesce the repeated txt.
121
+ e.g. vid [0 1 2 3 4 5 6 7 8] -> 3 splits -> [0 1 2] [3 4 5] [6 7 8]
122
+ txt [9 10]
123
+ repeat_concat ==> [0 1 2 9 10 3 4 5 9 10 6 7 8 9 10]
124
+ 1. argsort re-index ==> [0 1 2 3 4 5 6 7 8 9 9 9 10 10 10]
125
+ split ==> vid_out [0 1 2 3 4 5 6 7 8] txt_out [9 9 9 10 10 10]
126
+ 2. reshape & mean for each sample to coalesce the repeated txt.
127
+ """
128
+ vid_out, txt_out = all[src_idx].split([len(vid_idx), txt_idx_len])
129
+ txt_out_coalesced = []
130
+ for txt, repeat_time in zip(txt_out.split(repeat_txt_len), txt_repeat_list):
131
+ txt = txt.reshape(-1, repeat_time, *txt.shape[1:]).mean(1)
132
+ txt_out_coalesced.append(txt)
133
+ return vid_out, torch.cat(txt_out_coalesced)
134
+
135
+ # Note: Backward of torch.index_select is non-deterministic when existing repeated index,
136
+ # the difference may cumulative like torch.repeat_interleave, so we use vanilla index here.
137
+ return (
138
+ lambda vid, txt: torch.cat([vid, txt])[tgt_idx],
139
+ lambda all: unconcat_coalesce(all),
140
+ )
141
+
142
+
143
+ def rearrange(
144
+ hid: torch.FloatTensor, # (L c)
145
+ hid_shape: torch.LongTensor, # (b n)
146
+ pattern: str,
147
+ **kwargs: Dict[str, int],
148
+ ) -> Tuple[
149
+ torch.FloatTensor,
150
+ torch.LongTensor,
151
+ ]:
152
+ return flatten([einops.rearrange(h, pattern, **kwargs) for h in unflatten(hid, hid_shape)])
153
+
154
+
155
+ def rearrange_idx(
156
+ hid_shape: torch.LongTensor, # (b n)
157
+ pattern: str,
158
+ **kwargs: Dict[str, int],
159
+ ) -> Tuple[Callable, Callable, torch.LongTensor]:
160
+ hid_idx = torch.arange(hid_shape.prod(-1).sum(), device=hid_shape.device).unsqueeze(-1)
161
+ tgt_idx, tgt_shape = rearrange(hid_idx, hid_shape, pattern, **kwargs)
162
+ tgt_idx = tgt_idx.squeeze(-1)
163
+ src_idx = torch.argsort(tgt_idx)
164
+ return (
165
+ lambda hid: torch.index_select(hid, 0, tgt_idx),
166
+ lambda hid: torch.index_select(hid, 0, src_idx),
167
+ tgt_shape,
168
+ )
169
+
170
+
171
+ def repeat(
172
+ hid: torch.FloatTensor, # (L c)
173
+ hid_shape: torch.LongTensor, # (b n)
174
+ pattern: str,
175
+ **kwargs: Dict[str, torch.LongTensor], # (b)
176
+ ) -> Tuple[
177
+ torch.FloatTensor,
178
+ torch.LongTensor,
179
+ ]:
180
+ hid = unflatten(hid, hid_shape)
181
+ kwargs = [{k: v[i].item() for k, v in kwargs.items()} for i in range(len(hid))]
182
+ return flatten([einops.repeat(h, pattern, **a) for h, a in zip(hid, kwargs)])
183
+
184
+
185
+ def pack(
186
+ samples: List[torch.Tensor], # List of (h w c).
187
+ ) -> Tuple[
188
+ List[torch.Tensor], # groups [(b1 h1 w1 c1), (b2 h2 w2 c2)]
189
+ List[List[int]], # reversal indices.
190
+ ]:
191
+ batches = {}
192
+ indices = {}
193
+ for i, sample in enumerate(samples):
194
+ shape = sample.shape
195
+ batches[shape] = batches.get(shape, [])
196
+ indices[shape] = indices.get(shape, [])
197
+ batches[shape].append(sample)
198
+ indices[shape].append(i)
199
+
200
+ batches = list(map(torch.stack, batches.values()))
201
+ indices = list(indices.values())
202
+ return batches, indices
203
+
204
+
205
+ def unpack(
206
+ batches: List[torch.Tensor],
207
+ indices: List[List[int]],
208
+ ) -> List[torch.Tensor]:
209
+ samples = [None] * (max(chain(*indices)) + 1)
210
+ for batch, index in zip(batches, indices):
211
+ for sample, i in zip(batch.unbind(), index):
212
+ samples[i] = sample
213
+ return samples
214
+
215
+
216
+ def window(
217
+ hid: torch.FloatTensor, # (L c)
218
+ hid_shape: torch.LongTensor, # (b n)
219
+ window_fn: Callable[[torch.Tensor], List[torch.Tensor]],
220
+ ):
221
+ hid = unflatten(hid, hid_shape)
222
+ hid = list(map(window_fn, hid))
223
+ hid_windows = torch.tensor(list(map(len, hid)), device=hid_shape.device)
224
+ hid, hid_shape = flatten(list(chain(*hid)))
225
+ return hid, hid_shape, hid_windows
226
+
227
+
228
+ def window_idx(
229
+ hid_shape: torch.LongTensor, # (b n)
230
+ window_fn: Callable[[torch.Tensor], List[torch.Tensor]],
231
+ ):
232
+ hid_idx = torch.arange(hid_shape.prod(-1).sum(), device=hid_shape.device).unsqueeze(-1)
233
+ tgt_idx, tgt_shape, tgt_windows = window(hid_idx, hid_shape, window_fn)
234
+ tgt_idx = tgt_idx.squeeze(-1)
235
+ src_idx = torch.argsort(tgt_idx)
236
+ return (
237
+ lambda hid: torch.index_select(hid, 0, tgt_idx),
238
+ lambda hid: torch.index_select(hid, 0, src_idx),
239
+ tgt_shape,
240
+ tgt_windows,
241
+ )
models/dit_v2/nablocks/__init__.py ADDED
@@ -0,0 +1,26 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # // Copyright (c) 2025 Bytedance Ltd. and/or its affiliates
2
+ # //
3
+ # // Licensed under the Apache License, Version 2.0 (the "License");
4
+ # // you may not use this file except in compliance with the License.
5
+ # // You may obtain a copy of the License at
6
+ # //
7
+ # // http://www.apache.org/licenses/LICENSE-2.0
8
+ # //
9
+ # // Unless required by applicable law or agreed to in writing, software
10
+ # // distributed under the License is distributed on an "AS IS" BASIS,
11
+ # // WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # // See the License for the specific language governing permissions and
13
+ # // limitations under the License.
14
+
15
+ from .mmsr_block import NaMMSRTransformerBlock
16
+
17
+
18
+ nadit_blocks = {
19
+ "mmdit_sr": NaMMSRTransformerBlock,
20
+ }
21
+
22
+
23
+ def get_nablock(block_type: str):
24
+ if block_type in nadit_blocks:
25
+ return nadit_blocks[block_type]
26
+ raise NotImplementedError(f"{block_type} is not supported")
models/dit_v2/nablocks/attention/__init__.py ADDED
@@ -0,0 +1,25 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # // Copyright (c) 2025 Bytedance Ltd. and/or its affiliates
2
+ # //
3
+ # // Licensed under the Apache License, Version 2.0 (the "License");
4
+ # // you may not use this file except in compliance with the License.
5
+ # // You may obtain a copy of the License at
6
+ # //
7
+ # // http://www.apache.org/licenses/LICENSE-2.0
8
+ # //
9
+ # // Unless required by applicable law or agreed to in writing, software
10
+ # // distributed under the License is distributed on an "AS IS" BASIS,
11
+ # // WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # // See the License for the specific language governing permissions and
13
+ # // limitations under the License.
14
+
15
+ from .mmattn import NaMMAttention
16
+
17
+ attns = {
18
+ "mm_full": NaMMAttention,
19
+ }
20
+
21
+
22
+ def get_attn(attn_type: str):
23
+ if attn_type in attns:
24
+ return attns[attn_type]
25
+ raise NotImplementedError(f"{attn_type} is not supported")
models/dit_v2/nablocks/attention/mmattn.py ADDED
@@ -0,0 +1,266 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # // Copyright (c) 2025 Bytedance Ltd. and/or its affiliates
2
+ # //
3
+ # // Licensed under the Apache License, Version 2.0 (the "License");
4
+ # // you may not use this file except in compliance with the License.
5
+ # // You may obtain a copy of the License at
6
+ # //
7
+ # // http://www.apache.org/licenses/LICENSE-2.0
8
+ # //
9
+ # // Unless required by applicable law or agreed to in writing, software
10
+ # // distributed under the License is distributed on an "AS IS" BASIS,
11
+ # // WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # // See the License for the specific language governing permissions and
13
+ # // limitations under the License.
14
+
15
+ from typing import Optional, Tuple, Union
16
+ import torch
17
+ from einops import rearrange
18
+ from torch import nn
19
+ from torch.nn import functional as F
20
+ from torch.nn.modules.utils import _triple
21
+
22
+ from common.cache import Cache
23
+ from common.distributed.ops import gather_heads_scatter_seq, gather_seq_scatter_heads_qkv
24
+
25
+ from ... import na
26
+ from ...attention import FlashAttentionVarlen
27
+ from ...mm import MMArg, MMModule
28
+ from ...normalization import norm_layer_type
29
+ from ...rope import get_na_rope
30
+ from ...window import get_window_op
31
+ from itertools import chain
32
+
33
+
34
+ class NaMMAttention(nn.Module):
35
+ def __init__(
36
+ self,
37
+ vid_dim: int,
38
+ txt_dim: int,
39
+ heads: int,
40
+ head_dim: int,
41
+ qk_bias: bool,
42
+ qk_norm: norm_layer_type,
43
+ qk_norm_eps: float,
44
+ rope_type: Optional[str],
45
+ rope_dim: int,
46
+ shared_weights: bool,
47
+ **kwargs,
48
+ ):
49
+ super().__init__()
50
+ dim = MMArg(vid_dim, txt_dim)
51
+ inner_dim = heads * head_dim
52
+ qkv_dim = inner_dim * 3
53
+ self.head_dim = head_dim
54
+ self.proj_qkv = MMModule(
55
+ nn.Linear, dim, qkv_dim, bias=qk_bias, shared_weights=shared_weights
56
+ )
57
+ self.proj_out = MMModule(nn.Linear, inner_dim, dim, shared_weights=shared_weights)
58
+ self.norm_q = MMModule(
59
+ qk_norm,
60
+ dim=head_dim,
61
+ eps=qk_norm_eps,
62
+ elementwise_affine=True,
63
+ shared_weights=shared_weights,
64
+ )
65
+ self.norm_k = MMModule(
66
+ qk_norm,
67
+ dim=head_dim,
68
+ eps=qk_norm_eps,
69
+ elementwise_affine=True,
70
+ shared_weights=shared_weights,
71
+ )
72
+
73
+ self.rope = get_na_rope(rope_type=rope_type, dim=rope_dim)
74
+ self.attn = FlashAttentionVarlen()
75
+
76
+ def forward(
77
+ self,
78
+ vid: torch.FloatTensor, # l c
79
+ txt: torch.FloatTensor, # l c
80
+ vid_shape: torch.LongTensor, # b 3
81
+ txt_shape: torch.LongTensor, # b 1
82
+ cache: Cache,
83
+ ) -> Tuple[
84
+ torch.FloatTensor,
85
+ torch.FloatTensor,
86
+ ]:
87
+ vid_qkv, txt_qkv = self.proj_qkv(vid, txt)
88
+ vid_qkv = gather_seq_scatter_heads_qkv(
89
+ vid_qkv,
90
+ seq_dim=0,
91
+ qkv_shape=vid_shape,
92
+ cache=cache.namespace("vid"),
93
+ )
94
+ txt_qkv = gather_seq_scatter_heads_qkv(
95
+ txt_qkv,
96
+ seq_dim=0,
97
+ qkv_shape=txt_shape,
98
+ cache=cache.namespace("txt"),
99
+ )
100
+ vid_qkv = rearrange(vid_qkv, "l (o h d) -> l o h d", o=3, d=self.head_dim)
101
+ txt_qkv = rearrange(txt_qkv, "l (o h d) -> l o h d", o=3, d=self.head_dim)
102
+
103
+ vid_q, vid_k, vid_v = vid_qkv.unbind(1)
104
+ txt_q, txt_k, txt_v = txt_qkv.unbind(1)
105
+
106
+ vid_q, txt_q = self.norm_q(vid_q, txt_q)
107
+ vid_k, txt_k = self.norm_k(vid_k, txt_k)
108
+
109
+ if self.rope:
110
+ if self.rope.mm:
111
+ vid_q, vid_k, txt_q, txt_k = self.rope(
112
+ vid_q, vid_k, vid_shape, txt_q, txt_k, txt_shape, cache
113
+ )
114
+ else:
115
+ vid_q, vid_k = self.rope(vid_q, vid_k, vid_shape, cache)
116
+
117
+ vid_len = cache("vid_len", lambda: vid_shape.prod(-1))
118
+ txt_len = cache("txt_len", lambda: txt_shape.prod(-1))
119
+ all_len = cache("all_len", lambda: vid_len + txt_len)
120
+
121
+ concat, unconcat = cache("mm_pnp", lambda: na.concat_idx(vid_len, txt_len))
122
+
123
+ attn = self.attn(
124
+ q=concat(vid_q, txt_q).bfloat16(),
125
+ k=concat(vid_k, txt_k).bfloat16(),
126
+ v=concat(vid_v, txt_v).bfloat16(),
127
+ cu_seqlens_q=cache("mm_seqlens", lambda: F.pad(all_len.cumsum(0), (1, 0)).int()),
128
+ cu_seqlens_k=cache("mm_seqlens", lambda: F.pad(all_len.cumsum(0), (1, 0)).int()),
129
+ max_seqlen_q=cache("mm_maxlen", lambda: all_len.max().item()),
130
+ max_seqlen_k=cache("mm_maxlen", lambda: all_len.max().item()),
131
+ ).type_as(vid_q)
132
+
133
+ attn = rearrange(attn, "l h d -> l (h d)")
134
+ vid_out, txt_out = unconcat(attn)
135
+ vid_out = gather_heads_scatter_seq(vid_out, head_dim=1, seq_dim=0)
136
+ txt_out = gather_heads_scatter_seq(txt_out, head_dim=1, seq_dim=0)
137
+
138
+ vid_out, txt_out = self.proj_out(vid_out, txt_out)
139
+ return vid_out, txt_out
140
+
141
+
142
+ class NaSwinAttention(NaMMAttention):
143
+ def __init__(
144
+ self,
145
+ *args,
146
+ window: Union[int, Tuple[int, int, int]],
147
+ window_method: str,
148
+ **kwargs,
149
+ ):
150
+ super().__init__(*args, **kwargs)
151
+ self.window = _triple(window)
152
+ self.window_method = window_method
153
+ assert all(map(lambda v: isinstance(v, int) and v >= 0, self.window))
154
+
155
+ self.window_op = get_window_op(window_method)
156
+
157
+ def forward(
158
+ self,
159
+ vid: torch.FloatTensor, # l c
160
+ txt: torch.FloatTensor, # l c
161
+ vid_shape: torch.LongTensor, # b 3
162
+ txt_shape: torch.LongTensor, # b 1
163
+ cache: Cache,
164
+ ) -> Tuple[
165
+ torch.FloatTensor,
166
+ torch.FloatTensor,
167
+ ]:
168
+
169
+ vid_qkv, txt_qkv = self.proj_qkv(vid, txt)
170
+ vid_qkv = gather_seq_scatter_heads_qkv(
171
+ vid_qkv,
172
+ seq_dim=0,
173
+ qkv_shape=vid_shape,
174
+ cache=cache.namespace("vid"),
175
+ )
176
+ txt_qkv = gather_seq_scatter_heads_qkv(
177
+ txt_qkv,
178
+ seq_dim=0,
179
+ qkv_shape=txt_shape,
180
+ cache=cache.namespace("txt"),
181
+ )
182
+
183
+ # re-org the input seq for window attn
184
+ cache_win = cache.namespace(f"{self.window_method}_{self.window}_sd3")
185
+
186
+ def make_window(x: torch.Tensor):
187
+ t, h, w, _ = x.shape
188
+ window_slices = self.window_op((t, h, w), self.window)
189
+ return [x[st, sh, sw] for (st, sh, sw) in window_slices]
190
+
191
+ window_partition, window_reverse, window_shape, window_count = cache_win(
192
+ "win_transform",
193
+ lambda: na.window_idx(vid_shape, make_window),
194
+ )
195
+ vid_qkv_win = window_partition(vid_qkv)
196
+
197
+ vid_qkv_win = rearrange(vid_qkv_win, "l (o h d) -> l o h d", o=3, d=self.head_dim)
198
+ txt_qkv = rearrange(txt_qkv, "l (o h d) -> l o h d", o=3, d=self.head_dim)
199
+
200
+ vid_q, vid_k, vid_v = vid_qkv_win.unbind(1)
201
+ txt_q, txt_k, txt_v = txt_qkv.unbind(1)
202
+
203
+ vid_q, txt_q = self.norm_q(vid_q, txt_q)
204
+ vid_k, txt_k = self.norm_k(vid_k, txt_k)
205
+
206
+ txt_len = cache("txt_len", lambda: txt_shape.prod(-1))
207
+
208
+ vid_len_win = cache_win("vid_len", lambda: window_shape.prod(-1))
209
+ txt_len_win = cache_win("txt_len", lambda: txt_len.repeat_interleave(window_count))
210
+ all_len_win = cache_win("all_len", lambda: vid_len_win + txt_len_win)
211
+ concat_win, unconcat_win = cache_win(
212
+ "mm_pnp", lambda: na.repeat_concat_idx(vid_len_win, txt_len, window_count)
213
+ )
214
+
215
+ # window rope
216
+ if self.rope:
217
+ if self.rope.mm:
218
+ # repeat text q and k for window mmrope
219
+ _, num_h, _ = txt_q.shape
220
+ txt_q_repeat = rearrange(txt_q, "l h d -> l (h d)")
221
+ txt_q_repeat = na.unflatten(txt_q_repeat, txt_shape)
222
+ txt_q_repeat = [[x] * n for x, n in zip(txt_q_repeat, window_count)]
223
+ txt_q_repeat = list(chain(*txt_q_repeat))
224
+ txt_q_repeat, txt_shape_repeat = na.flatten(txt_q_repeat)
225
+ txt_q_repeat = rearrange(txt_q_repeat, "l (h d) -> l h d", h=num_h)
226
+
227
+ txt_k_repeat = rearrange(txt_k, "l h d -> l (h d)")
228
+ txt_k_repeat = na.unflatten(txt_k_repeat, txt_shape)
229
+ txt_k_repeat = [[x] * n for x, n in zip(txt_k_repeat, window_count)]
230
+ txt_k_repeat = list(chain(*txt_k_repeat))
231
+ txt_k_repeat, _ = na.flatten(txt_k_repeat)
232
+ txt_k_repeat = rearrange(txt_k_repeat, "l (h d) -> l h d", h=num_h)
233
+
234
+ vid_q, vid_k, txt_q, txt_k = self.rope(
235
+ vid_q, vid_k, window_shape, txt_q_repeat, txt_k_repeat, txt_shape_repeat, cache_win
236
+ )
237
+ else:
238
+ vid_q, vid_k = self.rope(vid_q, vid_k, window_shape, cache_win)
239
+
240
+ out = self.attn(
241
+ q=concat_win(vid_q, txt_q).bfloat16(),
242
+ k=concat_win(vid_k, txt_k).bfloat16(),
243
+ v=concat_win(vid_v, txt_v).bfloat16(),
244
+ cu_seqlens_q=cache_win(
245
+ "vid_seqlens_q", lambda: F.pad(all_len_win.cumsum(0), (1, 0)).int()
246
+ ),
247
+ cu_seqlens_k=cache_win(
248
+ "vid_seqlens_k", lambda: F.pad(all_len_win.cumsum(0), (1, 0)).int()
249
+ ),
250
+ max_seqlen_q=cache_win("vid_max_seqlen_q", lambda: all_len_win.max().item()),
251
+ max_seqlen_k=cache_win("vid_max_seqlen_k", lambda: all_len_win.max().item()),
252
+ ).type_as(vid_q)
253
+
254
+ # text pooling
255
+ vid_out, txt_out = unconcat_win(out)
256
+
257
+ vid_out = rearrange(vid_out, "l h d -> l (h d)")
258
+ txt_out = rearrange(txt_out, "l h d -> l (h d)")
259
+ vid_out = window_reverse(vid_out)
260
+
261
+ vid_out = gather_heads_scatter_seq(vid_out, head_dim=1, seq_dim=0)
262
+ txt_out = gather_heads_scatter_seq(txt_out, head_dim=1, seq_dim=0)
263
+
264
+ vid_out, txt_out = self.proj_out(vid_out, txt_out)
265
+
266
+ return vid_out, txt_out
models/dit_v2/nablocks/mmsr_block.py ADDED
@@ -0,0 +1,119 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # // Copyright (c) 2025 Bytedance Ltd. and/or its affiliates
2
+ # //
3
+ # // Licensed under the Apache License, Version 2.0 (the "License");
4
+ # // you may not use this file except in compliance with the License.
5
+ # // You may obtain a copy of the License at
6
+ # //
7
+ # // http://www.apache.org/licenses/LICENSE-2.0
8
+ # //
9
+ # // Unless required by applicable law or agreed to in writing, software
10
+ # // distributed under the License is distributed on an "AS IS" BASIS,
11
+ # // WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # // See the License for the specific language governing permissions and
13
+ # // limitations under the License.
14
+
15
+ from typing import Tuple
16
+ import torch
17
+ import torch.nn as nn
18
+
19
+ # from ..cache import Cache
20
+ from common.cache import Cache
21
+
22
+ from .attention.mmattn import NaSwinAttention
23
+ from ..mm import MMArg
24
+ from ..modulation import ada_layer_type
25
+ from ..normalization import norm_layer_type
26
+ from ..mm import MMArg, MMModule
27
+ from ..mlp import get_mlp
28
+
29
+
30
+ class NaMMSRTransformerBlock(nn.Module):
31
+ def __init__(
32
+ self,
33
+ *,
34
+ vid_dim: int,
35
+ txt_dim: int,
36
+ emb_dim: int,
37
+ heads: int,
38
+ head_dim: int,
39
+ expand_ratio: int,
40
+ norm: norm_layer_type,
41
+ norm_eps: float,
42
+ ada: ada_layer_type,
43
+ qk_bias: bool,
44
+ qk_norm: norm_layer_type,
45
+ mlp_type: str,
46
+ shared_weights: bool,
47
+ rope_type: str,
48
+ rope_dim: int,
49
+ is_last_layer: bool,
50
+ **kwargs,
51
+ ):
52
+ super().__init__()
53
+ dim = MMArg(vid_dim, txt_dim)
54
+ self.attn_norm = MMModule(norm, dim=dim, eps=norm_eps, elementwise_affine=False, shared_weights=shared_weights,)
55
+
56
+ self.attn = NaSwinAttention(
57
+ vid_dim=vid_dim,
58
+ txt_dim=txt_dim,
59
+ heads=heads,
60
+ head_dim=head_dim,
61
+ qk_bias=qk_bias,
62
+ qk_norm=qk_norm,
63
+ qk_norm_eps=norm_eps,
64
+ rope_type=rope_type,
65
+ rope_dim=rope_dim,
66
+ shared_weights=shared_weights,
67
+ window=kwargs.pop("window", None),
68
+ window_method=kwargs.pop("window_method", None),
69
+ )
70
+
71
+ self.mlp_norm = MMModule(norm, dim=dim, eps=norm_eps, elementwise_affine=False, shared_weights=shared_weights, vid_only=is_last_layer)
72
+ self.mlp = MMModule(
73
+ get_mlp(mlp_type),
74
+ dim=dim,
75
+ expand_ratio=expand_ratio,
76
+ shared_weights=shared_weights,
77
+ vid_only=is_last_layer
78
+ )
79
+ self.ada = MMModule(ada, dim=dim, emb_dim=emb_dim, layers=["attn", "mlp"], shared_weights=shared_weights, vid_only=is_last_layer)
80
+ self.is_last_layer = is_last_layer
81
+
82
+ def forward(
83
+ self,
84
+ vid: torch.FloatTensor, # l c
85
+ txt: torch.FloatTensor, # l c
86
+ vid_shape: torch.LongTensor, # b 3
87
+ txt_shape: torch.LongTensor, # b 1
88
+ emb: torch.FloatTensor,
89
+ cache: Cache,
90
+ ) -> Tuple[
91
+ torch.FloatTensor,
92
+ torch.FloatTensor,
93
+ torch.LongTensor,
94
+ torch.LongTensor,
95
+ ]:
96
+ hid_len = MMArg(
97
+ cache("vid_len", lambda: vid_shape.prod(-1)),
98
+ cache("txt_len", lambda: txt_shape.prod(-1)),
99
+ )
100
+ ada_kwargs = {
101
+ "emb": emb,
102
+ "hid_len": hid_len,
103
+ "cache": cache,
104
+ "branch_tag": MMArg("vid", "txt"),
105
+ }
106
+
107
+ vid_attn, txt_attn = self.attn_norm(vid, txt)
108
+ vid_attn, txt_attn = self.ada(vid_attn, txt_attn, layer="attn", mode="in", **ada_kwargs)
109
+ vid_attn, txt_attn = self.attn(vid_attn, txt_attn, vid_shape, txt_shape, cache)
110
+ vid_attn, txt_attn = self.ada(vid_attn, txt_attn, layer="attn", mode="out", **ada_kwargs)
111
+ vid_attn, txt_attn = (vid_attn + vid), (txt_attn + txt)
112
+
113
+ vid_mlp, txt_mlp = self.mlp_norm(vid_attn, txt_attn)
114
+ vid_mlp, txt_mlp = self.ada(vid_mlp, txt_mlp, layer="mlp", mode="in", **ada_kwargs)
115
+ vid_mlp, txt_mlp = self.mlp(vid_mlp, txt_mlp)
116
+ vid_mlp, txt_mlp = self.ada(vid_mlp, txt_mlp, layer="mlp", mode="out", **ada_kwargs)
117
+ vid_mlp, txt_mlp = (vid_mlp + vid_attn), (txt_mlp + txt_attn)
118
+
119
+ return vid_mlp, txt_mlp, vid_shape, txt_shape
models/dit_v2/nadit.py ADDED
@@ -0,0 +1,246 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # // Copyright (c) 2025 Bytedance Ltd. and/or its affiliates
2
+ # //
3
+ # // Licensed under the Apache License, Version 2.0 (the "License");
4
+ # // you may not use this file except in compliance with the License.
5
+ # // You may obtain a copy of the License at
6
+ # //
7
+ # // http://www.apache.org/licenses/LICENSE-2.0
8
+ # //
9
+ # // Unless required by applicable law or agreed to in writing, software
10
+ # // distributed under the License is distributed on an "AS IS" BASIS,
11
+ # // WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # // See the License for the specific language governing permissions and
13
+ # // limitations under the License.
14
+
15
+ from dataclasses import dataclass
16
+ from typing import List, Optional, Tuple, Union, Callable
17
+ import torch
18
+ from torch import nn
19
+
20
+ from common.cache import Cache
21
+ from common.distributed.ops import slice_inputs
22
+
23
+ from . import na
24
+ from .embedding import TimeEmbedding
25
+ from .modulation import get_ada_layer
26
+ from .nablocks import get_nablock
27
+ from .normalization import get_norm_layer
28
+ from .patch import get_na_patch_layers
29
+
30
+ # Fake func, no checkpointing is required for inference
31
+ def gradient_checkpointing(module: Union[Callable, nn.Module], *args, enabled: bool, **kwargs):
32
+ return module(*args, **kwargs)
33
+
34
+ @dataclass
35
+ class NaDiTOutput:
36
+ vid_sample: torch.Tensor
37
+
38
+
39
+ class NaDiT(nn.Module):
40
+ """
41
+ Native Resolution Diffusion Transformer (NaDiT)
42
+ """
43
+
44
+ gradient_checkpointing = False
45
+
46
+ def __init__(
47
+ self,
48
+ vid_in_channels: int,
49
+ vid_out_channels: int,
50
+ vid_dim: int,
51
+ txt_in_dim: Union[int, List[int]],
52
+ txt_dim: Optional[int],
53
+ emb_dim: int,
54
+ heads: int,
55
+ head_dim: int,
56
+ expand_ratio: int,
57
+ norm: Optional[str],
58
+ norm_eps: float,
59
+ ada: str,
60
+ qk_bias: bool,
61
+ qk_norm: Optional[str],
62
+ patch_size: Union[int, Tuple[int, int, int]],
63
+ num_layers: int,
64
+ block_type: Union[str, Tuple[str]],
65
+ mm_layers: Union[int, Tuple[bool]],
66
+ mlp_type: str = "normal",
67
+ patch_type: str = "v1",
68
+ rope_type: Optional[str] = "rope3d",
69
+ rope_dim: Optional[int] = None,
70
+ window: Optional[Tuple] = None,
71
+ window_method: Optional[Tuple[str]] = None,
72
+ msa_type: Optional[Tuple[str]] = None,
73
+ mca_type: Optional[Tuple[str]] = None,
74
+ txt_in_norm: Optional[str] = None,
75
+ txt_in_norm_scale_factor: int = 0.01,
76
+ txt_proj_type: Optional[str] = "linear",
77
+ vid_out_norm: Optional[str] = None,
78
+ **kwargs,
79
+ ):
80
+ ada = get_ada_layer(ada)
81
+ norm = get_norm_layer(norm)
82
+ qk_norm = get_norm_layer(qk_norm)
83
+ rope_dim = rope_dim if rope_dim is not None else head_dim // 2
84
+ if isinstance(block_type, str):
85
+ block_type = [block_type] * num_layers
86
+ elif len(block_type) != num_layers:
87
+ raise ValueError("The ``block_type`` list should equal to ``num_layers``.")
88
+ super().__init__()
89
+ NaPatchIn, NaPatchOut = get_na_patch_layers(patch_type)
90
+ self.vid_in = NaPatchIn(
91
+ in_channels=vid_in_channels,
92
+ patch_size=patch_size,
93
+ dim=vid_dim,
94
+ )
95
+ if not isinstance(txt_in_dim, int):
96
+ self.txt_in = nn.ModuleList([])
97
+ for in_dim in txt_in_dim:
98
+ txt_norm_layer = get_norm_layer(txt_in_norm)(txt_dim, norm_eps, True)
99
+ if txt_proj_type == "linear":
100
+ txt_proj_layer = nn.Linear(in_dim, txt_dim)
101
+ else:
102
+ txt_proj_layer = nn.Sequential(
103
+ nn.Linear(in_dim, in_dim), nn.GELU("tanh"), nn.Linear(in_dim, txt_dim)
104
+ )
105
+ torch.nn.init.constant_(txt_norm_layer.weight, txt_in_norm_scale_factor)
106
+ self.txt_in.append(
107
+ nn.Sequential(
108
+ txt_proj_layer,
109
+ txt_norm_layer,
110
+ )
111
+ )
112
+ else:
113
+ self.txt_in = (
114
+ nn.Linear(txt_in_dim, txt_dim)
115
+ if txt_in_dim and txt_in_dim != txt_dim
116
+ else nn.Identity()
117
+ )
118
+ self.emb_in = TimeEmbedding(
119
+ sinusoidal_dim=256,
120
+ hidden_dim=max(vid_dim, txt_dim),
121
+ output_dim=emb_dim,
122
+ )
123
+
124
+ if window is None or isinstance(window[0], int):
125
+ window = [window] * num_layers
126
+ if window_method is None or isinstance(window_method, str):
127
+ window_method = [window_method] * num_layers
128
+
129
+ if msa_type is None or isinstance(msa_type, str):
130
+ msa_type = [msa_type] * num_layers
131
+ if mca_type is None or isinstance(mca_type, str):
132
+ mca_type = [mca_type] * num_layers
133
+
134
+ self.blocks = nn.ModuleList(
135
+ [
136
+ get_nablock(block_type[i])(
137
+ vid_dim=vid_dim,
138
+ txt_dim=txt_dim,
139
+ emb_dim=emb_dim,
140
+ heads=heads,
141
+ head_dim=head_dim,
142
+ expand_ratio=expand_ratio,
143
+ norm=norm,
144
+ norm_eps=norm_eps,
145
+ ada=ada,
146
+ qk_bias=qk_bias,
147
+ qk_norm=qk_norm,
148
+ shared_weights=not (
149
+ (i < mm_layers) if isinstance(mm_layers, int) else mm_layers[i]
150
+ ),
151
+ mlp_type=mlp_type,
152
+ window=window[i],
153
+ window_method=window_method[i],
154
+ msa_type=msa_type[i],
155
+ mca_type=mca_type[i],
156
+ rope_type=rope_type,
157
+ rope_dim=rope_dim,
158
+ is_last_layer=(i == num_layers - 1),
159
+ **kwargs,
160
+ )
161
+ for i in range(num_layers)
162
+ ]
163
+ )
164
+
165
+ self.vid_out_norm = None
166
+ if vid_out_norm is not None:
167
+ self.vid_out_norm = get_norm_layer(vid_out_norm)(
168
+ dim=vid_dim,
169
+ eps=norm_eps,
170
+ elementwise_affine=True,
171
+ )
172
+ self.vid_out_ada = ada(
173
+ dim=vid_dim,
174
+ emb_dim=emb_dim,
175
+ layers=["out"],
176
+ modes=["in"],
177
+ )
178
+
179
+ self.vid_out = NaPatchOut(
180
+ out_channels=vid_out_channels,
181
+ patch_size=patch_size,
182
+ dim=vid_dim,
183
+ )
184
+
185
+ def set_gradient_checkpointing(self, enable: bool):
186
+ self.gradient_checkpointing = enable
187
+
188
+ def forward(
189
+ self,
190
+ vid: torch.FloatTensor, # l c
191
+ txt: Union[torch.FloatTensor, List[torch.FloatTensor]], # l c
192
+ vid_shape: torch.LongTensor, # b 3
193
+ txt_shape: Union[torch.LongTensor, List[torch.LongTensor]], # b 1
194
+ timestep: Union[int, float, torch.IntTensor, torch.FloatTensor], # b
195
+ disable_cache: bool = False, # for test
196
+ ):
197
+ cache = Cache(disable=disable_cache)
198
+
199
+ # slice vid after patching in when using sequence parallelism
200
+ if isinstance(txt, list):
201
+ assert isinstance(self.txt_in, nn.ModuleList)
202
+ txt = [
203
+ na.unflatten(fc(i), s) for fc, i, s in zip(self.txt_in, txt, txt_shape)
204
+ ] # B L D
205
+ txt, txt_shape = na.flatten([torch.cat(t, dim=0) for t in zip(*txt)])
206
+ txt = slice_inputs(txt, dim=0)
207
+ else:
208
+ txt = slice_inputs(txt, dim=0)
209
+ txt = self.txt_in(txt)
210
+
211
+ # Video input.
212
+ # Sequence parallel slicing is done inside patching class.
213
+ vid, vid_shape = self.vid_in(vid, vid_shape, cache)
214
+
215
+ # Embedding input.
216
+ emb = self.emb_in(timestep, device=vid.device, dtype=vid.dtype)
217
+
218
+ # Body
219
+ for i, block in enumerate(self.blocks):
220
+ vid, txt, vid_shape, txt_shape = gradient_checkpointing(
221
+ enabled=(self.gradient_checkpointing and self.training),
222
+ module=block,
223
+ vid=vid,
224
+ txt=txt,
225
+ vid_shape=vid_shape,
226
+ txt_shape=txt_shape,
227
+ emb=emb,
228
+ cache=cache,
229
+ )
230
+
231
+ # Video output norm.
232
+ if self.vid_out_norm:
233
+ vid = self.vid_out_norm(vid)
234
+ vid = self.vid_out_ada(
235
+ vid,
236
+ emb=emb,
237
+ layer="out",
238
+ mode="in",
239
+ hid_len=cache("vid_len", lambda: vid_shape.prod(-1)),
240
+ cache=cache,
241
+ branch_tag="vid",
242
+ )
243
+
244
+ # Video output.
245
+ vid, vid_shape = self.vid_out(vid, vid_shape, cache)
246
+ return NaDiTOutput(vid_sample=vid)
models/dit_v2/normalization.py ADDED
@@ -0,0 +1,63 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # // Copyright (c) 2025 Bytedance Ltd. and/or its affiliates
2
+ # //
3
+ # // Licensed under the Apache License, Version 2.0 (the "License");
4
+ # // you may not use this file except in compliance with the License.
5
+ # // You may obtain a copy of the License at
6
+ # //
7
+ # // http://www.apache.org/licenses/LICENSE-2.0
8
+ # //
9
+ # // Unless required by applicable law or agreed to in writing, software
10
+ # // distributed under the License is distributed on an "AS IS" BASIS,
11
+ # // WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # // See the License for the specific language governing permissions and
13
+ # // limitations under the License.
14
+
15
+ from typing import Callable, Optional
16
+ from diffusers.models.normalization import RMSNorm
17
+ from torch import nn
18
+
19
+ # (dim: int, eps: float, elementwise_affine: bool)
20
+ norm_layer_type = Callable[[int, float, bool], nn.Module]
21
+
22
+
23
+ def get_norm_layer(norm_type: Optional[str]) -> norm_layer_type:
24
+
25
+ def _norm_layer(dim: int, eps: float, elementwise_affine: bool):
26
+ if norm_type is None:
27
+ return nn.Identity()
28
+
29
+ if norm_type == "layer":
30
+ return nn.LayerNorm(
31
+ normalized_shape=dim,
32
+ eps=eps,
33
+ elementwise_affine=elementwise_affine,
34
+ )
35
+
36
+ if norm_type == "rms":
37
+ return RMSNorm(
38
+ dim=dim,
39
+ eps=eps,
40
+ elementwise_affine=elementwise_affine,
41
+ )
42
+
43
+ if norm_type == "fusedln":
44
+ from apex.normalization import FusedLayerNorm
45
+
46
+ return FusedLayerNorm(
47
+ normalized_shape=dim,
48
+ elementwise_affine=elementwise_affine,
49
+ eps=eps,
50
+ )
51
+
52
+ if norm_type == "fusedrms":
53
+ from apex.normalization import FusedRMSNorm
54
+
55
+ return FusedRMSNorm(
56
+ normalized_shape=dim,
57
+ elementwise_affine=elementwise_affine,
58
+ eps=eps,
59
+ )
60
+
61
+ raise NotImplementedError(f"{norm_type} is not supported")
62
+
63
+ return _norm_layer
models/dit_v2/patch/__init__.py ADDED
@@ -0,0 +1,19 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # // Copyright (c) 2025 Bytedance Ltd. and/or its affiliates
2
+ # //
3
+ # // Licensed under the Apache License, Version 2.0 (the "License");
4
+ # // you may not use this file except in compliance with the License.
5
+ # // You may obtain a copy of the License at
6
+ # //
7
+ # // http://www.apache.org/licenses/LICENSE-2.0
8
+ # //
9
+ # // Unless required by applicable law or agreed to in writing, software
10
+ # // distributed under the License is distributed on an "AS IS" BASIS,
11
+ # // WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # // See the License for the specific language governing permissions and
13
+ # // limitations under the License.
14
+
15
+ def get_na_patch_layers(patch_type="v1"):
16
+ assert patch_type in ["v1"]
17
+ if patch_type == "v1":
18
+ from .patch_v1 import NaPatchIn, NaPatchOut
19
+ return NaPatchIn, NaPatchOut
models/dit_v2/patch/patch_v1.py ADDED
@@ -0,0 +1,127 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # // Copyright (c) 2025 Bytedance Ltd. and/or its affiliates
2
+ # //
3
+ # // Licensed under the Apache License, Version 2.0 (the "License");
4
+ # // you may not use this file except in compliance with the License.
5
+ # // You may obtain a copy of the License at
6
+ # //
7
+ # // http://www.apache.org/licenses/LICENSE-2.0
8
+ # //
9
+ # // Unless required by applicable law or agreed to in writing, software
10
+ # // distributed under the License is distributed on an "AS IS" BASIS,
11
+ # // WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # // See the License for the specific language governing permissions and
13
+ # // limitations under the License.
14
+
15
+ from typing import Tuple, Union
16
+ import torch
17
+ from einops import rearrange
18
+ from torch import nn
19
+ from torch.nn.modules.utils import _triple
20
+
21
+ from common.cache import Cache
22
+ from common.distributed.ops import gather_outputs, slice_inputs
23
+
24
+ from .. import na
25
+
26
+
27
+ class PatchIn(nn.Module):
28
+ def __init__(
29
+ self,
30
+ in_channels: int,
31
+ patch_size: Union[int, Tuple[int, int, int]],
32
+ dim: int,
33
+ ):
34
+ super().__init__()
35
+ t, h, w = _triple(patch_size)
36
+ self.patch_size = t, h, w
37
+ self.proj = nn.Linear(in_channels * t * h * w, dim)
38
+
39
+ def forward(
40
+ self,
41
+ vid: torch.Tensor,
42
+ ) -> torch.Tensor:
43
+ t, h, w = self.patch_size
44
+ if t > 1:
45
+ assert vid.size(2) % t == 1
46
+ vid = torch.cat([vid[:, :, :1]] * (t - 1) + [vid], dim=2)
47
+ vid = rearrange(vid, "b c (T t) (H h) (W w) -> b T H W (t h w c)", t=t, h=h, w=w)
48
+ vid = self.proj(vid)
49
+ return vid
50
+
51
+
52
+ class PatchOut(nn.Module):
53
+ def __init__(
54
+ self,
55
+ out_channels: int,
56
+ patch_size: Union[int, Tuple[int, int, int]],
57
+ dim: int,
58
+ ):
59
+ super().__init__()
60
+ t, h, w = _triple(patch_size)
61
+ self.patch_size = t, h, w
62
+ self.proj = nn.Linear(dim, out_channels * t * h * w)
63
+
64
+ def forward(
65
+ self,
66
+ vid: torch.Tensor,
67
+ ) -> torch.Tensor:
68
+ t, h, w = self.patch_size
69
+ vid = self.proj(vid)
70
+ vid = rearrange(vid, "b T H W (t h w c) -> b c (T t) (H h) (W w)", t=t, h=h, w=w)
71
+ if t > 1:
72
+ vid = vid[:, :, (t - 1) :]
73
+ return vid
74
+
75
+
76
+ class NaPatchIn(PatchIn):
77
+ def forward(
78
+ self,
79
+ vid: torch.Tensor, # l c
80
+ vid_shape: torch.LongTensor,
81
+ cache: Cache = Cache(disable=True), # for test
82
+ ) -> torch.Tensor:
83
+ cache = cache.namespace("patch")
84
+ vid_shape_before_patchify = cache("vid_shape_before_patchify", lambda: vid_shape)
85
+ t, h, w = self.patch_size
86
+ if not (t == h == w == 1):
87
+ vid = na.unflatten(vid, vid_shape)
88
+ for i in range(len(vid)):
89
+ if t > 1 and vid_shape_before_patchify[i, 0] % t != 0:
90
+ vid[i] = torch.cat([vid[i][:1]] * (t - vid[i].size(0) % t) + [vid[i]], dim=0)
91
+ vid[i] = rearrange(vid[i], "(T t) (H h) (W w) c -> T H W (t h w c)", t=t, h=h, w=w)
92
+ vid, vid_shape = na.flatten(vid)
93
+
94
+ # slice vid after patching in when using sequence parallelism
95
+ vid = slice_inputs(vid, dim=0)
96
+ vid = self.proj(vid)
97
+ return vid, vid_shape
98
+
99
+
100
+ class NaPatchOut(PatchOut):
101
+ def forward(
102
+ self,
103
+ vid: torch.FloatTensor, # l c
104
+ vid_shape: torch.LongTensor,
105
+ cache: Cache = Cache(disable=True), # for test
106
+ ) -> Tuple[
107
+ torch.FloatTensor,
108
+ torch.LongTensor,
109
+ ]:
110
+ cache = cache.namespace("patch")
111
+ vid_shape_before_patchify = cache.get("vid_shape_before_patchify")
112
+
113
+ t, h, w = self.patch_size
114
+ vid = self.proj(vid)
115
+ # gather vid before patching out when enabling sequence parallelism
116
+ vid = gather_outputs(
117
+ vid, gather_dim=0, padding_dim=0, unpad_shape=vid_shape, cache=cache.namespace("vid")
118
+ )
119
+ if not (t == h == w == 1):
120
+ vid = na.unflatten(vid, vid_shape)
121
+ for i in range(len(vid)):
122
+ vid[i] = rearrange(vid[i], "T H W (t h w c) -> (T t) (H h) (W w) c", t=t, h=h, w=w)
123
+ if t > 1 and vid_shape_before_patchify[i, 0] % t != 0:
124
+ vid[i] = vid[i][(t - vid_shape_before_patchify[i, 0] % t) :]
125
+ vid, vid_shape = na.flatten(vid)
126
+
127
+ return vid, vid_shape
models/dit_v2/rope.py ADDED
@@ -0,0 +1,150 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # // Copyright (c) 2025 Bytedance Ltd. and/or its affiliates
2
+ # //
3
+ # // Licensed under the Apache License, Version 2.0 (the "License");
4
+ # // you may not use this file except in compliance with the License.
5
+ # // You may obtain a copy of the License at
6
+ # //
7
+ # // http://www.apache.org/licenses/LICENSE-2.0
8
+ # //
9
+ # // Unless required by applicable law or agreed to in writing, software
10
+ # // distributed under the License is distributed on an "AS IS" BASIS,
11
+ # // WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # // See the License for the specific language governing permissions and
13
+ # // limitations under the License.
14
+
15
+ from functools import lru_cache
16
+ from typing import Optional, Tuple
17
+ import torch
18
+ from einops import rearrange
19
+ from rotary_embedding_torch import RotaryEmbedding, apply_rotary_emb
20
+ from torch import nn
21
+
22
+ from common.cache import Cache
23
+
24
+
25
+ class RotaryEmbeddingBase(nn.Module):
26
+ def __init__(self, dim: int, rope_dim: int):
27
+ super().__init__()
28
+ self.rope = RotaryEmbedding(
29
+ dim=dim // rope_dim,
30
+ freqs_for="pixel",
31
+ max_freq=256,
32
+ )
33
+ # 1. Set model.requires_grad_(True) after model creation will make
34
+ # the `requires_grad=False` for rope freqs no longer hold.
35
+ # 2. Even if we don't set requires_grad_(True) explicitly,
36
+ # FSDP is not memory efficient when handling fsdp_wrap
37
+ # with mixed requires_grad=True/False.
38
+ # With above consideration, it is easier just remove the freqs
39
+ # out of nn.Parameters when `learned_freq=False`
40
+ freqs = self.rope.freqs
41
+ del self.rope.freqs
42
+ self.rope.register_buffer("freqs", freqs.data)
43
+
44
+ @lru_cache(maxsize=128)
45
+ def get_axial_freqs(self, *dims):
46
+ return self.rope.get_axial_freqs(*dims)
47
+
48
+
49
+ class RotaryEmbedding3d(RotaryEmbeddingBase):
50
+ def __init__(self, dim: int):
51
+ super().__init__(dim, rope_dim=3)
52
+ self.mm = False
53
+
54
+ def forward(
55
+ self,
56
+ q: torch.FloatTensor, # b h l d
57
+ k: torch.FloatTensor, # b h l d
58
+ size: Tuple[int, int, int],
59
+ ) -> Tuple[
60
+ torch.FloatTensor,
61
+ torch.FloatTensor,
62
+ ]:
63
+ T, H, W = size
64
+ freqs = self.get_axial_freqs(T, H, W)
65
+ q = rearrange(q, "b h (T H W) d -> b h T H W d", T=T, H=H, W=W)
66
+ k = rearrange(k, "b h (T H W) d -> b h T H W d", T=T, H=H, W=W)
67
+ q = apply_rotary_emb(freqs, q.float()).to(q.dtype)
68
+ k = apply_rotary_emb(freqs, k.float()).to(k.dtype)
69
+ q = rearrange(q, "b h T H W d -> b h (T H W) d")
70
+ k = rearrange(k, "b h T H W d -> b h (T H W) d")
71
+ return q, k
72
+
73
+
74
+ class MMRotaryEmbeddingBase(RotaryEmbeddingBase):
75
+ def __init__(self, dim: int, rope_dim: int):
76
+ super().__init__(dim, rope_dim)
77
+ self.rope = RotaryEmbedding(
78
+ dim=dim // rope_dim,
79
+ freqs_for="lang",
80
+ theta=10000,
81
+ )
82
+ freqs = self.rope.freqs
83
+ del self.rope.freqs
84
+ self.rope.register_buffer("freqs", freqs.data)
85
+ self.mm = True
86
+
87
+
88
+ class NaMMRotaryEmbedding3d(MMRotaryEmbeddingBase):
89
+ def __init__(self, dim: int):
90
+ super().__init__(dim, rope_dim=3)
91
+
92
+ def forward(
93
+ self,
94
+ vid_q: torch.FloatTensor, # L h d
95
+ vid_k: torch.FloatTensor, # L h d
96
+ vid_shape: torch.LongTensor, # B 3
97
+ txt_q: torch.FloatTensor, # L h d
98
+ txt_k: torch.FloatTensor, # L h d
99
+ txt_shape: torch.LongTensor, # B 1
100
+ cache: Cache,
101
+ ) -> Tuple[
102
+ torch.FloatTensor,
103
+ torch.FloatTensor,
104
+ torch.FloatTensor,
105
+ torch.FloatTensor,
106
+ ]:
107
+ vid_freqs, txt_freqs = cache(
108
+ "mmrope_freqs_3d",
109
+ lambda: self.get_freqs(vid_shape, txt_shape),
110
+ )
111
+ vid_q = rearrange(vid_q, "L h d -> h L d")
112
+ vid_k = rearrange(vid_k, "L h d -> h L d")
113
+ vid_q = apply_rotary_emb(vid_freqs, vid_q.float()).to(vid_q.dtype)
114
+ vid_k = apply_rotary_emb(vid_freqs, vid_k.float()).to(vid_k.dtype)
115
+ vid_q = rearrange(vid_q, "h L d -> L h d")
116
+ vid_k = rearrange(vid_k, "h L d -> L h d")
117
+
118
+ txt_q = rearrange(txt_q, "L h d -> h L d")
119
+ txt_k = rearrange(txt_k, "L h d -> h L d")
120
+ txt_q = apply_rotary_emb(txt_freqs, txt_q.float()).to(txt_q.dtype)
121
+ txt_k = apply_rotary_emb(txt_freqs, txt_k.float()).to(txt_k.dtype)
122
+ txt_q = rearrange(txt_q, "h L d -> L h d")
123
+ txt_k = rearrange(txt_k, "h L d -> L h d")
124
+ return vid_q, vid_k, txt_q, txt_k
125
+
126
+ def get_freqs(
127
+ self,
128
+ vid_shape: torch.LongTensor,
129
+ txt_shape: torch.LongTensor,
130
+ ) -> Tuple[
131
+ torch.Tensor,
132
+ torch.Tensor,
133
+ ]:
134
+ vid_freqs = self.get_axial_freqs(1024, 128, 128)
135
+ txt_freqs = self.get_axial_freqs(1024)
136
+ vid_freq_list, txt_freq_list = [], []
137
+ for (f, h, w), l in zip(vid_shape.tolist(), txt_shape[:, 0].tolist()):
138
+ vid_freq = vid_freqs[l : l + f, :h, :w].reshape(-1, vid_freqs.size(-1))
139
+ txt_freq = txt_freqs[:l].repeat(1, 3).reshape(-1, vid_freqs.size(-1))
140
+ vid_freq_list.append(vid_freq)
141
+ txt_freq_list.append(txt_freq)
142
+ return torch.cat(vid_freq_list, dim=0), torch.cat(txt_freq_list, dim=0)
143
+
144
+
145
+ def get_na_rope(rope_type: Optional[str], dim: int):
146
+ if rope_type is None:
147
+ return None
148
+ if rope_type == "mmrope3d":
149
+ return NaMMRotaryEmbedding3d(dim=dim)
150
+ raise NotImplementedError(f"{rope_type} is not supported.")
models/dit_v2/window.py ADDED
@@ -0,0 +1,83 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # // Copyright (c) 2025 Bytedance Ltd. and/or its affiliates
2
+ # //
3
+ # // Licensed under the Apache License, Version 2.0 (the "License");
4
+ # // you may not use this file except in compliance with the License.
5
+ # // You may obtain a copy of the License at
6
+ # //
7
+ # // http://www.apache.org/licenses/LICENSE-2.0
8
+ # //
9
+ # // Unless required by applicable law or agreed to in writing, software
10
+ # // distributed under the License is distributed on an "AS IS" BASIS,
11
+ # // WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # // See the License for the specific language governing permissions and
13
+ # // limitations under the License.
14
+
15
+ from math import ceil
16
+ from typing import Tuple
17
+ import math
18
+
19
+ def get_window_op(name: str):
20
+ if name == "720pwin_by_size_bysize":
21
+ return make_720Pwindows_bysize
22
+ if name == "720pswin_by_size_bysize":
23
+ return make_shifted_720Pwindows_bysize
24
+ raise ValueError(f"Unknown windowing method: {name}")
25
+
26
+
27
+ # -------------------------------- Windowing -------------------------------- #
28
+ def make_720Pwindows_bysize(size: Tuple[int, int, int], num_windows: Tuple[int, int, int]):
29
+ t, h, w = size
30
+ resized_nt, resized_nh, resized_nw = num_windows
31
+ #cal windows under 720p
32
+ scale = math.sqrt((45 * 80) / (h * w))
33
+ resized_h, resized_w = round(h * scale), round(w * scale)
34
+ wh, ww = ceil(resized_h / resized_nh), ceil(resized_w / resized_nw) # window size.
35
+ wt = ceil(min(t, 30) / resized_nt) # window size.
36
+ nt, nh, nw = ceil(t / wt), ceil(h / wh), ceil(w / ww) # window size.
37
+ return [
38
+ (
39
+ slice(it * wt, min((it + 1) * wt, t)),
40
+ slice(ih * wh, min((ih + 1) * wh, h)),
41
+ slice(iw * ww, min((iw + 1) * ww, w)),
42
+ )
43
+ for iw in range(nw)
44
+ if min((iw + 1) * ww, w) > iw * ww
45
+ for ih in range(nh)
46
+ if min((ih + 1) * wh, h) > ih * wh
47
+ for it in range(nt)
48
+ if min((it + 1) * wt, t) > it * wt
49
+ ]
50
+
51
+ def make_shifted_720Pwindows_bysize(size: Tuple[int, int, int], num_windows: Tuple[int, int, int]):
52
+ t, h, w = size
53
+ resized_nt, resized_nh, resized_nw = num_windows
54
+ #cal windows under 720p
55
+ scale = math.sqrt((45 * 80) / (h * w))
56
+ resized_h, resized_w = round(h * scale), round(w * scale)
57
+ wh, ww = ceil(resized_h / resized_nh), ceil(resized_w / resized_nw) # window size.
58
+ wt = ceil(min(t, 30) / resized_nt) # window size.
59
+
60
+ st, sh, sw = ( # shift size.
61
+ 0.5 if wt < t else 0,
62
+ 0.5 if wh < h else 0,
63
+ 0.5 if ww < w else 0,
64
+ )
65
+ nt, nh, nw = ceil((t - st) / wt), ceil((h - sh) / wh), ceil((w - sw) / ww) # window size.
66
+ nt, nh, nw = ( # number of window.
67
+ nt + 1 if st > 0 else 1,
68
+ nh + 1 if sh > 0 else 1,
69
+ nw + 1 if sw > 0 else 1,
70
+ )
71
+ return [
72
+ (
73
+ slice(max(int((it - st) * wt), 0), min(int((it - st + 1) * wt), t)),
74
+ slice(max(int((ih - sh) * wh), 0), min(int((ih - sh + 1) * wh), h)),
75
+ slice(max(int((iw - sw) * ww), 0), min(int((iw - sw + 1) * ww), w)),
76
+ )
77
+ for iw in range(nw)
78
+ if min(int((iw - sw + 1) * ww), w) > max(int((iw - sw) * ww), 0)
79
+ for ih in range(nh)
80
+ if min(int((ih - sh + 1) * wh), h) > max(int((ih - sh) * wh), 0)
81
+ for it in range(nt)
82
+ if min(int((it - st + 1) * wt), t) > max(int((it - st) * wt), 0)
83
+ ]
models/video_vae_v3/modules/attn_video_vae.py ADDED
@@ -0,0 +1,1345 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) 2023 HuggingFace Team
2
+ # Copyright (c) 2025 ByteDance Ltd. and/or its affiliates.
3
+ # SPDX-License-Identifier: Apache License, Version 2.0 (the "License")
4
+ #
5
+ # This file has been modified by ByteDance Ltd. and/or its affiliates. on 1st June 2025
6
+ #
7
+ # Original file was released under Apache License, Version 2.0 (the "License"), with the full license text
8
+ # available at http://www.apache.org/licenses/LICENSE-2.0.
9
+ #
10
+ # This modified file is released under the same license.
11
+
12
+
13
+ from contextlib import nullcontext
14
+ from typing import Literal, Optional, Tuple, Union
15
+ import diffusers
16
+ import torch
17
+ import torch.nn as nn
18
+ import torch.nn.functional as F
19
+ from diffusers.models.attention_processor import Attention, SpatialNorm
20
+ from diffusers.models.autoencoders.vae import DecoderOutput, DiagonalGaussianDistribution
21
+ from diffusers.models.downsampling import Downsample2D
22
+ from diffusers.models.lora import LoRACompatibleConv
23
+ from diffusers.models.modeling_outputs import AutoencoderKLOutput
24
+ from diffusers.models.resnet import ResnetBlock2D
25
+ from diffusers.models.unets.unet_2d_blocks import DownEncoderBlock2D, UpDecoderBlock2D
26
+ from diffusers.models.upsampling import Upsample2D
27
+ from diffusers.utils import is_torch_version
28
+ from diffusers.utils.accelerate_utils import apply_forward_hook
29
+ from einops import rearrange
30
+
31
+ from common.distributed.advanced import get_sequence_parallel_world_size
32
+ from common.logger import get_logger
33
+ from models.video_vae_v3.modules.causal_inflation_lib import (
34
+ InflatedCausalConv3d,
35
+ causal_norm_wrapper,
36
+ init_causal_conv3d,
37
+ remove_head,
38
+ )
39
+ from models.video_vae_v3.modules.context_parallel_lib import (
40
+ causal_conv_gather_outputs,
41
+ causal_conv_slice_inputs,
42
+ )
43
+ from models.video_vae_v3.modules.global_config import set_norm_limit
44
+ from models.video_vae_v3.modules.types import (
45
+ CausalAutoencoderOutput,
46
+ CausalDecoderOutput,
47
+ CausalEncoderOutput,
48
+ MemoryState,
49
+ _inflation_mode_t,
50
+ _memory_device_t,
51
+ _receptive_field_t,
52
+ )
53
+
54
+ logger = get_logger(__name__) # pylint: disable=invalid-name
55
+
56
+
57
+ class Upsample3D(Upsample2D):
58
+ """A 3D upsampling layer with an optional convolution."""
59
+
60
+ def __init__(
61
+ self,
62
+ *args,
63
+ inflation_mode: _inflation_mode_t = "tail",
64
+ temporal_up: bool = False,
65
+ spatial_up: bool = True,
66
+ slicing: bool = False,
67
+ **kwargs,
68
+ ):
69
+ super().__init__(*args, **kwargs)
70
+ conv = self.conv if self.name == "conv" else self.Conv2d_0
71
+
72
+ assert type(conv) is not nn.ConvTranspose2d
73
+ # Note: lora_layer is not passed into constructor in the original implementation.
74
+ # So we make a simplification.
75
+ conv = init_causal_conv3d(
76
+ self.channels,
77
+ self.out_channels,
78
+ 3,
79
+ padding=1,
80
+ inflation_mode=inflation_mode,
81
+ )
82
+
83
+ self.temporal_up = temporal_up
84
+ self.spatial_up = spatial_up
85
+ self.temporal_ratio = 2 if temporal_up else 1
86
+ self.spatial_ratio = 2 if spatial_up else 1
87
+ self.slicing = slicing
88
+
89
+ assert not self.interpolate
90
+ # [Override] MAGViT v2 implementation
91
+ if not self.interpolate:
92
+ upscale_ratio = (self.spatial_ratio**2) * self.temporal_ratio
93
+ self.upscale_conv = nn.Conv3d(
94
+ self.channels, self.channels * upscale_ratio, kernel_size=1, padding=0
95
+ )
96
+ identity = (
97
+ torch.eye(self.channels)
98
+ .repeat(upscale_ratio, 1)
99
+ .reshape_as(self.upscale_conv.weight)
100
+ )
101
+ self.upscale_conv.weight.data.copy_(identity)
102
+ nn.init.zeros_(self.upscale_conv.bias)
103
+
104
+ if self.name == "conv":
105
+ self.conv = conv
106
+ else:
107
+ self.Conv2d_0 = conv
108
+
109
+ def forward(
110
+ self,
111
+ hidden_states: torch.FloatTensor,
112
+ output_size: Optional[int] = None,
113
+ memory_state: MemoryState = MemoryState.DISABLED,
114
+ **kwargs,
115
+ ) -> torch.FloatTensor:
116
+ assert hidden_states.shape[1] == self.channels
117
+
118
+ if hasattr(self, "norm") and self.norm is not None:
119
+ # [Overridden] change to causal norm.
120
+ hidden_states = causal_norm_wrapper(self.norm, hidden_states)
121
+
122
+ if self.use_conv_transpose:
123
+ return self.conv(hidden_states)
124
+
125
+ if self.slicing:
126
+ split_size = hidden_states.size(2) // 2
127
+ hidden_states = list(
128
+ hidden_states.split([split_size, hidden_states.size(2) - split_size], dim=2)
129
+ )
130
+ else:
131
+ hidden_states = [hidden_states]
132
+
133
+ for i in range(len(hidden_states)):
134
+ hidden_states[i] = self.upscale_conv(hidden_states[i])
135
+ hidden_states[i] = rearrange(
136
+ hidden_states[i],
137
+ "b (x y z c) f h w -> b c (f z) (h x) (w y)",
138
+ x=self.spatial_ratio,
139
+ y=self.spatial_ratio,
140
+ z=self.temporal_ratio,
141
+ )
142
+
143
+ # [Overridden] For causal temporal conv
144
+ if self.temporal_up and memory_state != MemoryState.ACTIVE:
145
+ hidden_states[0] = remove_head(hidden_states[0])
146
+
147
+ if not self.slicing:
148
+ hidden_states = hidden_states[0]
149
+
150
+ if self.use_conv:
151
+ if self.name == "conv":
152
+ hidden_states = self.conv(hidden_states, memory_state=memory_state)
153
+ else:
154
+ hidden_states = self.Conv2d_0(hidden_states, memory_state=memory_state)
155
+
156
+ if not self.slicing:
157
+ return hidden_states
158
+ else:
159
+ return torch.cat(hidden_states, dim=2)
160
+
161
+
162
+ class Downsample3D(Downsample2D):
163
+ """A 3D downsampling layer with an optional convolution."""
164
+
165
+ def __init__(
166
+ self,
167
+ *args,
168
+ inflation_mode: _inflation_mode_t = "tail",
169
+ spatial_down: bool = False,
170
+ temporal_down: bool = False,
171
+ **kwargs,
172
+ ):
173
+ super().__init__(*args, **kwargs)
174
+ conv = self.conv
175
+ self.temporal_down = temporal_down
176
+ self.spatial_down = spatial_down
177
+
178
+ self.temporal_ratio = 2 if temporal_down else 1
179
+ self.spatial_ratio = 2 if spatial_down else 1
180
+
181
+ self.temporal_kernel = 3 if temporal_down else 1
182
+ self.spatial_kernel = 3 if spatial_down else 1
183
+
184
+ if type(conv) in [nn.Conv2d, LoRACompatibleConv]:
185
+ # Note: lora_layer is not passed into constructor in the original implementation.
186
+ # So we make a simplification.
187
+ conv = init_causal_conv3d(
188
+ self.channels,
189
+ self.out_channels,
190
+ kernel_size=(self.temporal_kernel, self.spatial_kernel, self.spatial_kernel),
191
+ stride=(self.temporal_ratio, self.spatial_ratio, self.spatial_ratio),
192
+ padding=(
193
+ 1 if self.temporal_down else 0,
194
+ self.padding if self.spatial_down else 0,
195
+ self.padding if self.spatial_down else 0,
196
+ ),
197
+ inflation_mode=inflation_mode,
198
+ )
199
+ elif type(conv) is nn.AvgPool2d:
200
+ assert self.channels == self.out_channels
201
+ conv = nn.AvgPool3d(
202
+ kernel_size=(self.temporal_ratio, self.spatial_ratio, self.spatial_ratio),
203
+ stride=(self.temporal_ratio, self.spatial_ratio, self.spatial_ratio),
204
+ )
205
+ else:
206
+ raise NotImplementedError
207
+
208
+ if self.name == "conv":
209
+ self.Conv2d_0 = conv
210
+ self.conv = conv
211
+ else:
212
+ self.conv = conv
213
+
214
+ def forward(
215
+ self,
216
+ hidden_states: torch.FloatTensor,
217
+ memory_state: MemoryState = MemoryState.DISABLED,
218
+ **kwargs,
219
+ ) -> torch.FloatTensor:
220
+
221
+ assert hidden_states.shape[1] == self.channels
222
+
223
+ if hasattr(self, "norm") and self.norm is not None:
224
+ # [Overridden] change to causal norm.
225
+ hidden_states = causal_norm_wrapper(self.norm, hidden_states)
226
+
227
+ if self.use_conv and self.padding == 0 and self.spatial_down:
228
+ pad = (0, 1, 0, 1)
229
+ hidden_states = F.pad(hidden_states, pad, mode="constant", value=0)
230
+
231
+ assert hidden_states.shape[1] == self.channels
232
+
233
+ hidden_states = self.conv(hidden_states, memory_state=memory_state)
234
+
235
+ return hidden_states
236
+
237
+
238
+ class ResnetBlock3D(ResnetBlock2D):
239
+ def __init__(
240
+ self,
241
+ *args,
242
+ inflation_mode: _inflation_mode_t = "tail",
243
+ time_receptive_field: _receptive_field_t = "half",
244
+ slicing: bool = False,
245
+ **kwargs,
246
+ ):
247
+ super().__init__(*args, **kwargs)
248
+ self.conv1 = init_causal_conv3d(
249
+ self.in_channels,
250
+ self.out_channels,
251
+ kernel_size=(1, 3, 3) if time_receptive_field == "half" else (3, 3, 3),
252
+ stride=1,
253
+ padding=(0, 1, 1) if time_receptive_field == "half" else (1, 1, 1),
254
+ inflation_mode=inflation_mode,
255
+ )
256
+
257
+ self.conv2 = init_causal_conv3d(
258
+ self.out_channels,
259
+ self.conv2.out_channels,
260
+ kernel_size=3,
261
+ stride=1,
262
+ padding=1,
263
+ inflation_mode=inflation_mode,
264
+ )
265
+
266
+ if self.up:
267
+ assert type(self.upsample) is Upsample2D
268
+ self.upsample = Upsample3D(
269
+ self.in_channels,
270
+ use_conv=False,
271
+ inflation_mode=inflation_mode,
272
+ slicing=slicing,
273
+ )
274
+ elif self.down:
275
+ assert type(self.downsample) is Downsample2D
276
+ self.downsample = Downsample3D(
277
+ self.in_channels,
278
+ use_conv=False,
279
+ padding=1,
280
+ name="op",
281
+ inflation_mode=inflation_mode,
282
+ )
283
+
284
+ if self.use_in_shortcut:
285
+ self.conv_shortcut = init_causal_conv3d(
286
+ self.in_channels,
287
+ self.conv_shortcut.out_channels,
288
+ kernel_size=1,
289
+ stride=1,
290
+ padding=0,
291
+ bias=(self.conv_shortcut.bias is not None),
292
+ inflation_mode=inflation_mode,
293
+ )
294
+
295
+ def forward(
296
+ self, input_tensor, temb, memory_state: MemoryState = MemoryState.DISABLED, **kwargs
297
+ ):
298
+ hidden_states = input_tensor
299
+
300
+ hidden_states = causal_norm_wrapper(self.norm1, hidden_states)
301
+
302
+ hidden_states = self.nonlinearity(hidden_states)
303
+
304
+ if self.upsample is not None:
305
+ # upsample_nearest_nhwc fails with large batch sizes.
306
+ # see https://github.com/huggingface/diffusers/issues/984
307
+ if hidden_states.shape[0] >= 64:
308
+ input_tensor = input_tensor.contiguous()
309
+ hidden_states = hidden_states.contiguous()
310
+ input_tensor = self.upsample(input_tensor, memory_state=memory_state)
311
+ hidden_states = self.upsample(hidden_states, memory_state=memory_state)
312
+ elif self.downsample is not None:
313
+ input_tensor = self.downsample(input_tensor, memory_state=memory_state)
314
+ hidden_states = self.downsample(hidden_states, memory_state=memory_state)
315
+
316
+ hidden_states = self.conv1(hidden_states, memory_state=memory_state)
317
+
318
+ if self.time_emb_proj is not None:
319
+ if not self.skip_time_act:
320
+ temb = self.nonlinearity(temb)
321
+ temb = self.time_emb_proj(temb)[:, :, None, None]
322
+
323
+ if temb is not None and self.time_embedding_norm == "default":
324
+ hidden_states = hidden_states + temb
325
+
326
+ hidden_states = causal_norm_wrapper(self.norm2, hidden_states)
327
+
328
+ if temb is not None and self.time_embedding_norm == "scale_shift":
329
+ scale, shift = torch.chunk(temb, 2, dim=1)
330
+ hidden_states = hidden_states * (1 + scale) + shift
331
+
332
+ hidden_states = self.nonlinearity(hidden_states)
333
+
334
+ hidden_states = self.dropout(hidden_states)
335
+ hidden_states = self.conv2(hidden_states, memory_state=memory_state)
336
+
337
+ if self.conv_shortcut is not None:
338
+ input_tensor = self.conv_shortcut(input_tensor, memory_state=memory_state)
339
+
340
+ output_tensor = (input_tensor + hidden_states) / self.output_scale_factor
341
+
342
+ return output_tensor
343
+
344
+
345
+ class DownEncoderBlock3D(DownEncoderBlock2D):
346
+ def __init__(
347
+ self,
348
+ in_channels: int,
349
+ out_channels: int,
350
+ dropout: float = 0.0,
351
+ num_layers: int = 1,
352
+ resnet_eps: float = 1e-6,
353
+ resnet_time_scale_shift: str = "default",
354
+ resnet_act_fn: str = "swish",
355
+ resnet_groups: int = 32,
356
+ resnet_pre_norm: bool = True,
357
+ output_scale_factor: float = 1.0,
358
+ add_downsample: bool = True,
359
+ downsample_padding: int = 1,
360
+ inflation_mode: _inflation_mode_t = "tail",
361
+ time_receptive_field: _receptive_field_t = "half",
362
+ temporal_down: bool = True,
363
+ spatial_down: bool = True,
364
+ ):
365
+ super().__init__(
366
+ in_channels=in_channels,
367
+ out_channels=out_channels,
368
+ dropout=dropout,
369
+ num_layers=num_layers,
370
+ resnet_eps=resnet_eps,
371
+ resnet_time_scale_shift=resnet_time_scale_shift,
372
+ resnet_act_fn=resnet_act_fn,
373
+ resnet_groups=resnet_groups,
374
+ resnet_pre_norm=resnet_pre_norm,
375
+ output_scale_factor=output_scale_factor,
376
+ add_downsample=add_downsample,
377
+ downsample_padding=downsample_padding,
378
+ )
379
+ resnets = []
380
+ temporal_modules = []
381
+
382
+ for i in range(num_layers):
383
+ in_channels = in_channels if i == 0 else out_channels
384
+ resnets.append(
385
+ # [Override] Replace module.
386
+ ResnetBlock3D(
387
+ in_channels=in_channels,
388
+ out_channels=out_channels,
389
+ temb_channels=None,
390
+ eps=resnet_eps,
391
+ groups=resnet_groups,
392
+ dropout=dropout,
393
+ time_embedding_norm=resnet_time_scale_shift,
394
+ non_linearity=resnet_act_fn,
395
+ output_scale_factor=output_scale_factor,
396
+ pre_norm=resnet_pre_norm,
397
+ inflation_mode=inflation_mode,
398
+ time_receptive_field=time_receptive_field,
399
+ )
400
+ )
401
+ temporal_modules.append(nn.Identity())
402
+
403
+ self.resnets = nn.ModuleList(resnets)
404
+ self.temporal_modules = nn.ModuleList(temporal_modules)
405
+
406
+ if add_downsample:
407
+ self.downsamplers = nn.ModuleList(
408
+ [
409
+ # [Override] Replace module.
410
+ Downsample3D(
411
+ out_channels,
412
+ use_conv=True,
413
+ out_channels=out_channels,
414
+ padding=downsample_padding,
415
+ name="op",
416
+ temporal_down=temporal_down,
417
+ spatial_down=spatial_down,
418
+ inflation_mode=inflation_mode,
419
+ )
420
+ ]
421
+ )
422
+ else:
423
+ self.downsamplers = None
424
+
425
+ def forward(
426
+ self,
427
+ hidden_states: torch.FloatTensor,
428
+ memory_state: MemoryState = MemoryState.DISABLED,
429
+ **kwargs,
430
+ ) -> torch.FloatTensor:
431
+ for resnet, temporal in zip(self.resnets, self.temporal_modules):
432
+ hidden_states = resnet(hidden_states, temb=None, memory_state=memory_state)
433
+ hidden_states = temporal(hidden_states)
434
+
435
+ if self.downsamplers is not None:
436
+ for downsampler in self.downsamplers:
437
+ hidden_states = downsampler(hidden_states, memory_state=memory_state)
438
+
439
+ return hidden_states
440
+
441
+
442
+ class UpDecoderBlock3D(UpDecoderBlock2D):
443
+ def __init__(
444
+ self,
445
+ in_channels: int,
446
+ out_channels: int,
447
+ dropout: float = 0.0,
448
+ num_layers: int = 1,
449
+ resnet_eps: float = 1e-6,
450
+ resnet_time_scale_shift: str = "default", # default, spatial
451
+ resnet_act_fn: str = "swish",
452
+ resnet_groups: int = 32,
453
+ resnet_pre_norm: bool = True,
454
+ output_scale_factor: float = 1.0,
455
+ add_upsample: bool = True,
456
+ temb_channels: Optional[int] = None,
457
+ inflation_mode: _inflation_mode_t = "tail",
458
+ time_receptive_field: _receptive_field_t = "half",
459
+ temporal_up: bool = True,
460
+ spatial_up: bool = True,
461
+ slicing: bool = False,
462
+ ):
463
+ super().__init__(
464
+ in_channels=in_channels,
465
+ out_channels=out_channels,
466
+ dropout=dropout,
467
+ num_layers=num_layers,
468
+ resnet_eps=resnet_eps,
469
+ resnet_time_scale_shift=resnet_time_scale_shift,
470
+ resnet_act_fn=resnet_act_fn,
471
+ resnet_groups=resnet_groups,
472
+ resnet_pre_norm=resnet_pre_norm,
473
+ output_scale_factor=output_scale_factor,
474
+ add_upsample=add_upsample,
475
+ temb_channels=temb_channels,
476
+ )
477
+ resnets = []
478
+ temporal_modules = []
479
+
480
+ for i in range(num_layers):
481
+ input_channels = in_channels if i == 0 else out_channels
482
+
483
+ resnets.append(
484
+ # [Override] Replace module.
485
+ ResnetBlock3D(
486
+ in_channels=input_channels,
487
+ out_channels=out_channels,
488
+ temb_channels=temb_channels,
489
+ eps=resnet_eps,
490
+ groups=resnet_groups,
491
+ dropout=dropout,
492
+ time_embedding_norm=resnet_time_scale_shift,
493
+ non_linearity=resnet_act_fn,
494
+ output_scale_factor=output_scale_factor,
495
+ pre_norm=resnet_pre_norm,
496
+ inflation_mode=inflation_mode,
497
+ time_receptive_field=time_receptive_field,
498
+ slicing=slicing,
499
+ )
500
+ )
501
+
502
+ temporal_modules.append(nn.Identity())
503
+
504
+ self.resnets = nn.ModuleList(resnets)
505
+ self.temporal_modules = nn.ModuleList(temporal_modules)
506
+
507
+ if add_upsample:
508
+ # [Override] Replace module & use learnable upsample
509
+ self.upsamplers = nn.ModuleList(
510
+ [
511
+ Upsample3D(
512
+ out_channels,
513
+ use_conv=True,
514
+ out_channels=out_channels,
515
+ temporal_up=temporal_up,
516
+ spatial_up=spatial_up,
517
+ interpolate=False,
518
+ inflation_mode=inflation_mode,
519
+ slicing=slicing,
520
+ )
521
+ ]
522
+ )
523
+ else:
524
+ self.upsamplers = None
525
+
526
+ def forward(
527
+ self,
528
+ hidden_states: torch.FloatTensor,
529
+ temb: Optional[torch.FloatTensor] = None,
530
+ memory_state: MemoryState = MemoryState.DISABLED,
531
+ ) -> torch.FloatTensor:
532
+ for resnet, temporal in zip(self.resnets, self.temporal_modules):
533
+ hidden_states = resnet(hidden_states, temb=None, memory_state=memory_state)
534
+ hidden_states = temporal(hidden_states)
535
+
536
+ if self.upsamplers is not None:
537
+ for upsampler in self.upsamplers:
538
+ hidden_states = upsampler(hidden_states, memory_state=memory_state)
539
+
540
+ return hidden_states
541
+
542
+
543
+ class UNetMidBlock3D(nn.Module):
544
+ def __init__(
545
+ self,
546
+ in_channels: int,
547
+ temb_channels: int,
548
+ dropout: float = 0.0,
549
+ num_layers: int = 1,
550
+ resnet_eps: float = 1e-6,
551
+ resnet_time_scale_shift: str = "default", # default, spatial
552
+ resnet_act_fn: str = "swish",
553
+ resnet_groups: int = 32,
554
+ resnet_pre_norm: bool = True,
555
+ add_attention: bool = True,
556
+ attention_head_dim: int = 1,
557
+ output_scale_factor: float = 1.0,
558
+ inflation_mode: _inflation_mode_t = "tail",
559
+ time_receptive_field: _receptive_field_t = "half",
560
+ ):
561
+ super().__init__()
562
+ resnet_groups = resnet_groups if resnet_groups is not None else min(in_channels // 4, 32)
563
+ self.add_attention = add_attention
564
+
565
+ # there is always at least one resnet
566
+ resnets = [
567
+ # [Override] Replace module.
568
+ ResnetBlock3D(
569
+ in_channels=in_channels,
570
+ out_channels=in_channels,
571
+ temb_channels=temb_channels,
572
+ eps=resnet_eps,
573
+ groups=resnet_groups,
574
+ dropout=dropout,
575
+ time_embedding_norm=resnet_time_scale_shift,
576
+ non_linearity=resnet_act_fn,
577
+ output_scale_factor=output_scale_factor,
578
+ pre_norm=resnet_pre_norm,
579
+ inflation_mode=inflation_mode,
580
+ time_receptive_field=time_receptive_field,
581
+ )
582
+ ]
583
+ attentions = []
584
+
585
+ if attention_head_dim is None:
586
+ logger.warn(
587
+ f"It is not recommend to pass `attention_head_dim=None`. "
588
+ f"Defaulting `attention_head_dim` to `in_channels`: {in_channels}."
589
+ )
590
+ attention_head_dim = in_channels
591
+
592
+ for _ in range(num_layers):
593
+ if self.add_attention:
594
+ attentions.append(
595
+ Attention(
596
+ in_channels,
597
+ heads=in_channels // attention_head_dim,
598
+ dim_head=attention_head_dim,
599
+ rescale_output_factor=output_scale_factor,
600
+ eps=resnet_eps,
601
+ norm_num_groups=(
602
+ resnet_groups if resnet_time_scale_shift == "default" else None
603
+ ),
604
+ spatial_norm_dim=(
605
+ temb_channels if resnet_time_scale_shift == "spatial" else None
606
+ ),
607
+ residual_connection=True,
608
+ bias=True,
609
+ upcast_softmax=True,
610
+ _from_deprecated_attn_block=True,
611
+ )
612
+ )
613
+ else:
614
+ attentions.append(None)
615
+
616
+ resnets.append(
617
+ ResnetBlock3D(
618
+ in_channels=in_channels,
619
+ out_channels=in_channels,
620
+ temb_channels=temb_channels,
621
+ eps=resnet_eps,
622
+ groups=resnet_groups,
623
+ dropout=dropout,
624
+ time_embedding_norm=resnet_time_scale_shift,
625
+ non_linearity=resnet_act_fn,
626
+ output_scale_factor=output_scale_factor,
627
+ pre_norm=resnet_pre_norm,
628
+ inflation_mode=inflation_mode,
629
+ time_receptive_field=time_receptive_field,
630
+ )
631
+ )
632
+
633
+ self.attentions = nn.ModuleList(attentions)
634
+ self.resnets = nn.ModuleList(resnets)
635
+
636
+ def forward(self, hidden_states, temb=None, memory_state: MemoryState = MemoryState.DISABLED):
637
+ video_length, frame_height, frame_width = hidden_states.size()[-3:]
638
+ hidden_states = self.resnets[0](hidden_states, temb, memory_state=memory_state)
639
+ for attn, resnet in zip(self.attentions, self.resnets[1:]):
640
+ if attn is not None:
641
+ hidden_states = rearrange(hidden_states, "b c f h w -> (b f) c h w")
642
+ hidden_states = attn(hidden_states, temb=temb)
643
+ hidden_states = rearrange(
644
+ hidden_states, "(b f) c h w -> b c f h w", f=video_length
645
+ )
646
+ hidden_states = resnet(hidden_states, temb, memory_state=memory_state)
647
+
648
+ return hidden_states
649
+
650
+
651
+ class Encoder3D(nn.Module):
652
+ r"""
653
+ [Override] override most logics to support extra condition input and causal conv
654
+
655
+ The `Encoder` layer of a variational autoencoder that encodes
656
+ its input into a latent representation.
657
+
658
+ Args:
659
+ in_channels (`int`, *optional*, defaults to 3):
660
+ The number of input channels.
661
+ out_channels (`int`, *optional*, defaults to 3):
662
+ The number of output channels.
663
+ down_block_types (`Tuple[str, ...]`, *optional*, defaults to `("DownEncoderBlock2D",)`):
664
+ The types of down blocks to use.
665
+ See `~diffusers.models.unet_2d_blocks.get_down_block`
666
+ for available options.
667
+ block_out_channels (`Tuple[int, ...]`, *optional*, defaults to `(64,)`):
668
+ The number of output channels for each block.
669
+ layers_per_block (`int`, *optional*, defaults to 2):
670
+ The number of layers per block.
671
+ norm_num_groups (`int`, *optional*, defaults to 32):
672
+ The number of groups for normalization.
673
+ act_fn (`str`, *optional*, defaults to `"silu"`):
674
+ The activation function to use.
675
+ See `~diffusers.models.activations.get_activation` for available options.
676
+ double_z (`bool`, *optional*, defaults to `True`):
677
+ Whether to double the number of output channels for the last block.
678
+ """
679
+
680
+ def __init__(
681
+ self,
682
+ in_channels: int = 3,
683
+ out_channels: int = 3,
684
+ down_block_types: Tuple[str, ...] = ("DownEncoderBlock3D",),
685
+ block_out_channels: Tuple[int, ...] = (64,),
686
+ layers_per_block: int = 2,
687
+ norm_num_groups: int = 32,
688
+ act_fn: str = "silu",
689
+ double_z: bool = True,
690
+ mid_block_add_attention=True,
691
+ # [Override] add extra_cond_dim, temporal down num
692
+ temporal_down_num: int = 2,
693
+ extra_cond_dim: int = None,
694
+ gradient_checkpoint: bool = False,
695
+ inflation_mode: _inflation_mode_t = "tail",
696
+ time_receptive_field: _receptive_field_t = "half",
697
+ ):
698
+ super().__init__()
699
+ self.layers_per_block = layers_per_block
700
+ self.temporal_down_num = temporal_down_num
701
+
702
+ self.conv_in = init_causal_conv3d(
703
+ in_channels,
704
+ block_out_channels[0],
705
+ kernel_size=3,
706
+ stride=1,
707
+ padding=1,
708
+ inflation_mode=inflation_mode,
709
+ )
710
+
711
+ self.mid_block = None
712
+ self.down_blocks = nn.ModuleList([])
713
+ self.extra_cond_dim = extra_cond_dim
714
+
715
+ self.conv_extra_cond = nn.ModuleList([])
716
+
717
+ # down
718
+ output_channel = block_out_channels[0]
719
+ for i, down_block_type in enumerate(down_block_types):
720
+ input_channel = output_channel
721
+ output_channel = block_out_channels[i]
722
+ is_final_block = i == len(block_out_channels) - 1
723
+ # [Override] to support temporal down block design
724
+ is_temporal_down_block = i >= len(block_out_channels) - self.temporal_down_num - 1
725
+ # Note: take the last ones
726
+
727
+ assert down_block_type == "DownEncoderBlock3D"
728
+
729
+ down_block = DownEncoderBlock3D(
730
+ num_layers=self.layers_per_block,
731
+ in_channels=input_channel,
732
+ out_channels=output_channel,
733
+ add_downsample=not is_final_block,
734
+ resnet_eps=1e-6,
735
+ downsample_padding=0,
736
+ # Note: Don't know why set it as 0
737
+ resnet_act_fn=act_fn,
738
+ resnet_groups=norm_num_groups,
739
+ temporal_down=is_temporal_down_block,
740
+ spatial_down=True,
741
+ inflation_mode=inflation_mode,
742
+ time_receptive_field=time_receptive_field,
743
+ )
744
+ self.down_blocks.append(down_block)
745
+
746
+ def zero_module(module):
747
+ # Zero out the parameters of a module and return it.
748
+ for p in module.parameters():
749
+ p.detach().zero_()
750
+ return module
751
+
752
+ self.conv_extra_cond.append(
753
+ zero_module(
754
+ nn.Conv3d(extra_cond_dim, output_channel, kernel_size=1, stride=1, padding=0)
755
+ )
756
+ if self.extra_cond_dim is not None and self.extra_cond_dim > 0
757
+ else None
758
+ )
759
+
760
+ # mid
761
+ self.mid_block = UNetMidBlock3D(
762
+ in_channels=block_out_channels[-1],
763
+ resnet_eps=1e-6,
764
+ resnet_act_fn=act_fn,
765
+ output_scale_factor=1,
766
+ resnet_time_scale_shift="default",
767
+ attention_head_dim=block_out_channels[-1],
768
+ resnet_groups=norm_num_groups,
769
+ temb_channels=None,
770
+ add_attention=mid_block_add_attention,
771
+ inflation_mode=inflation_mode,
772
+ time_receptive_field=time_receptive_field,
773
+ )
774
+
775
+ # out
776
+ self.conv_norm_out = nn.GroupNorm(
777
+ num_channels=block_out_channels[-1], num_groups=norm_num_groups, eps=1e-6
778
+ )
779
+ self.conv_act = nn.SiLU()
780
+
781
+ conv_out_channels = 2 * out_channels if double_z else out_channels
782
+ self.conv_out = init_causal_conv3d(
783
+ block_out_channels[-1], conv_out_channels, 3, padding=1, inflation_mode=inflation_mode
784
+ )
785
+
786
+ self.gradient_checkpointing = gradient_checkpoint
787
+
788
+ def forward(
789
+ self,
790
+ sample: torch.FloatTensor,
791
+ extra_cond=None,
792
+ memory_state: MemoryState = MemoryState.DISABLED,
793
+ ) -> torch.FloatTensor:
794
+ r"""The forward method of the `Encoder` class."""
795
+ sample = self.conv_in(sample, memory_state=memory_state)
796
+ if self.training and self.gradient_checkpointing:
797
+
798
+ def create_custom_forward(module):
799
+ def custom_forward(*inputs):
800
+ return module(*inputs)
801
+
802
+ return custom_forward
803
+
804
+ # down
805
+ # [Override] add extra block and extra cond
806
+ for down_block, extra_block in zip(self.down_blocks, self.conv_extra_cond):
807
+ sample = torch.utils.checkpoint.checkpoint(
808
+ create_custom_forward(down_block), sample, memory_state, use_reentrant=False
809
+ )
810
+ if extra_block is not None:
811
+ sample = sample + F.interpolate(extra_block(extra_cond), size=sample.shape[2:])
812
+
813
+ # middle
814
+ sample = self.mid_block(sample, memory_state=memory_state)
815
+
816
+ # sample = torch.utils.checkpoint.checkpoint(
817
+ # create_custom_forward(self.mid_block), sample, use_reentrant=False
818
+ # )
819
+
820
+ else:
821
+ # down
822
+ # [Override] add extra block and extra cond
823
+ for down_block, extra_block in zip(self.down_blocks, self.conv_extra_cond):
824
+ sample = down_block(sample, memory_state=memory_state)
825
+ if extra_block is not None:
826
+ sample = sample + F.interpolate(extra_block(extra_cond), size=sample.shape[2:])
827
+
828
+ # middle
829
+ sample = self.mid_block(sample, memory_state=memory_state)
830
+
831
+ # post-process
832
+ sample = causal_norm_wrapper(self.conv_norm_out, sample)
833
+ sample = self.conv_act(sample)
834
+ sample = self.conv_out(sample, memory_state=memory_state)
835
+
836
+ return sample
837
+
838
+
839
+ class Decoder3D(nn.Module):
840
+ r"""
841
+ The `Decoder` layer of a variational autoencoder that
842
+ decodes its latent representation into an output sample.
843
+
844
+ Args:
845
+ in_channels (`int`, *optional*, defaults to 3):
846
+ The number of input channels.
847
+ out_channels (`int`, *optional*, defaults to 3):
848
+ The number of output channels.
849
+ up_block_types (`Tuple[str, ...]`, *optional*, defaults to `("UpDecoderBlock2D",)`):
850
+ The types of up blocks to use.
851
+ See `~diffusers.models.unet_2d_blocks.get_up_block` for available options.
852
+ block_out_channels (`Tuple[int, ...]`, *optional*, defaults to `(64,)`):
853
+ The number of output channels for each block.
854
+ layers_per_block (`int`, *optional*, defaults to 2):
855
+ The number of layers per block.
856
+ norm_num_groups (`int`, *optional*, defaults to 32):
857
+ The number of groups for normalization.
858
+ act_fn (`str`, *optional*, defaults to `"silu"`):
859
+ The activation function to use.
860
+ See `~diffusers.models.activations.get_activation` for available options.
861
+ norm_type (`str`, *optional*, defaults to `"group"`):
862
+ The normalization type to use. Can be either `"group"` or `"spatial"`.
863
+ """
864
+
865
+ def __init__(
866
+ self,
867
+ in_channels: int = 3,
868
+ out_channels: int = 3,
869
+ up_block_types: Tuple[str, ...] = ("UpDecoderBlock3D",),
870
+ block_out_channels: Tuple[int, ...] = (64,),
871
+ layers_per_block: int = 2,
872
+ norm_num_groups: int = 32,
873
+ act_fn: str = "silu",
874
+ norm_type: str = "group", # group, spatial
875
+ mid_block_add_attention=True,
876
+ # [Override] add temporal up block
877
+ inflation_mode: _inflation_mode_t = "tail",
878
+ time_receptive_field: _receptive_field_t = "half",
879
+ temporal_up_num: int = 2,
880
+ slicing_up_num: int = 0,
881
+ gradient_checkpoint: bool = False,
882
+ ):
883
+ super().__init__()
884
+ self.layers_per_block = layers_per_block
885
+ self.temporal_up_num = temporal_up_num
886
+
887
+ self.conv_in = init_causal_conv3d(
888
+ in_channels,
889
+ block_out_channels[-1],
890
+ kernel_size=3,
891
+ stride=1,
892
+ padding=1,
893
+ inflation_mode=inflation_mode,
894
+ )
895
+
896
+ self.mid_block = None
897
+ self.up_blocks = nn.ModuleList([])
898
+
899
+ temb_channels = in_channels if norm_type == "spatial" else None
900
+
901
+ # mid
902
+ self.mid_block = UNetMidBlock3D(
903
+ in_channels=block_out_channels[-1],
904
+ resnet_eps=1e-6,
905
+ resnet_act_fn=act_fn,
906
+ output_scale_factor=1,
907
+ resnet_time_scale_shift="default" if norm_type == "group" else norm_type,
908
+ attention_head_dim=block_out_channels[-1],
909
+ resnet_groups=norm_num_groups,
910
+ temb_channels=temb_channels,
911
+ add_attention=mid_block_add_attention,
912
+ inflation_mode=inflation_mode,
913
+ time_receptive_field=time_receptive_field,
914
+ )
915
+
916
+ # up
917
+ reversed_block_out_channels = list(reversed(block_out_channels))
918
+ output_channel = reversed_block_out_channels[0]
919
+ print(f"slicing_up_num: {slicing_up_num}")
920
+ for i, up_block_type in enumerate(up_block_types):
921
+ prev_output_channel = output_channel
922
+ output_channel = reversed_block_out_channels[i]
923
+
924
+ is_final_block = i == len(block_out_channels) - 1
925
+ is_temporal_up_block = i < self.temporal_up_num
926
+ is_slicing_up_block = i >= len(block_out_channels) - slicing_up_num
927
+ # Note: Keep symmetric
928
+
929
+ assert up_block_type == "UpDecoderBlock3D"
930
+ up_block = UpDecoderBlock3D(
931
+ num_layers=self.layers_per_block + 1,
932
+ in_channels=prev_output_channel,
933
+ out_channels=output_channel,
934
+ add_upsample=not is_final_block,
935
+ resnet_eps=1e-6,
936
+ resnet_act_fn=act_fn,
937
+ resnet_groups=norm_num_groups,
938
+ resnet_time_scale_shift=norm_type,
939
+ temb_channels=temb_channels,
940
+ temporal_up=is_temporal_up_block,
941
+ slicing=is_slicing_up_block,
942
+ inflation_mode=inflation_mode,
943
+ time_receptive_field=time_receptive_field,
944
+ )
945
+ self.up_blocks.append(up_block)
946
+ prev_output_channel = output_channel
947
+
948
+ # out
949
+ if norm_type == "spatial":
950
+ self.conv_norm_out = SpatialNorm(block_out_channels[0], temb_channels)
951
+ else:
952
+ self.conv_norm_out = nn.GroupNorm(
953
+ num_channels=block_out_channels[0], num_groups=norm_num_groups, eps=1e-6
954
+ )
955
+ self.conv_act = nn.SiLU()
956
+ self.conv_out = init_causal_conv3d(
957
+ block_out_channels[0], out_channels, 3, padding=1, inflation_mode=inflation_mode
958
+ )
959
+
960
+ self.gradient_checkpointing = gradient_checkpoint
961
+
962
+ # Note: Just copy from Decoder.
963
+ def forward(
964
+ self,
965
+ sample: torch.FloatTensor,
966
+ latent_embeds: Optional[torch.FloatTensor] = None,
967
+ memory_state: MemoryState = MemoryState.DISABLED,
968
+ ) -> torch.FloatTensor:
969
+ r"""The forward method of the `Decoder` class."""
970
+
971
+ sample = self.conv_in(sample, memory_state=memory_state)
972
+
973
+ upscale_dtype = next(iter(self.up_blocks.parameters())).dtype
974
+ if self.training and self.gradient_checkpointing:
975
+
976
+ def create_custom_forward(module):
977
+ def custom_forward(*inputs):
978
+ return module(*inputs)
979
+
980
+ return custom_forward
981
+
982
+ if is_torch_version(">=", "1.11.0"):
983
+ sample = self.mid_block(sample, latent_embeds, memory_state=memory_state)
984
+ sample = sample.to(upscale_dtype)
985
+
986
+ # up
987
+ for up_block in self.up_blocks:
988
+ sample = torch.utils.checkpoint.checkpoint(
989
+ create_custom_forward(up_block),
990
+ sample,
991
+ latent_embeds,
992
+ memory_state,
993
+ use_reentrant=False,
994
+ )
995
+ else:
996
+ # middle
997
+ sample = self.mid_block(sample, latent_embeds, memory_state=memory_state)
998
+ sample = sample.to(upscale_dtype)
999
+
1000
+ # up
1001
+ for up_block in self.up_blocks:
1002
+ sample = torch.utils.checkpoint.checkpoint(
1003
+ create_custom_forward(up_block), sample, latent_embeds, memory_state
1004
+ )
1005
+ else:
1006
+ # middle
1007
+ sample = self.mid_block(sample, latent_embeds, memory_state=memory_state)
1008
+ sample = sample.to(upscale_dtype)
1009
+
1010
+ # up
1011
+ for up_block in self.up_blocks:
1012
+ sample = up_block(sample, latent_embeds, memory_state=memory_state)
1013
+
1014
+ # post-process
1015
+ sample = causal_norm_wrapper(self.conv_norm_out, sample)
1016
+ sample = self.conv_act(sample)
1017
+ sample = self.conv_out(sample, memory_state=memory_state)
1018
+
1019
+ return sample
1020
+
1021
+
1022
+ class AutoencoderKL(diffusers.AutoencoderKL):
1023
+ """
1024
+ We simply inherit the model code from diffusers
1025
+ """
1026
+
1027
+ def __init__(self, attention: bool = True, *args, **kwargs):
1028
+ super().__init__(*args, **kwargs)
1029
+
1030
+ # A hacky way to remove attention.
1031
+ if not attention:
1032
+ self.encoder.mid_block.attentions = torch.nn.ModuleList([None])
1033
+ self.decoder.mid_block.attentions = torch.nn.ModuleList([None])
1034
+
1035
+ def load_state_dict(self, state_dict, strict=True):
1036
+ # Newer version of diffusers changed the model keys,
1037
+ # causing incompatibility with old checkpoints.
1038
+ # They provided a method for conversion. We call conversion before loading state_dict.
1039
+ convert_deprecated_attention_blocks = getattr(
1040
+ self, "_convert_deprecated_attention_blocks", None
1041
+ )
1042
+ if callable(convert_deprecated_attention_blocks):
1043
+ convert_deprecated_attention_blocks(state_dict)
1044
+ return super().load_state_dict(state_dict, strict)
1045
+
1046
+
1047
+ class VideoAutoencoderKL(diffusers.AutoencoderKL):
1048
+ """
1049
+ We simply inherit the model code from diffusers
1050
+ """
1051
+
1052
+ def __init__(
1053
+ self,
1054
+ in_channels: int = 3,
1055
+ out_channels: int = 3,
1056
+ down_block_types: Tuple[str] = ("DownEncoderBlock3D",),
1057
+ up_block_types: Tuple[str] = ("UpDecoderBlock3D",),
1058
+ block_out_channels: Tuple[int] = (64,),
1059
+ layers_per_block: int = 1,
1060
+ act_fn: str = "silu",
1061
+ latent_channels: int = 4,
1062
+ norm_num_groups: int = 32,
1063
+ sample_size: int = 32,
1064
+ scaling_factor: float = 0.18215,
1065
+ force_upcast: float = True,
1066
+ attention: bool = True,
1067
+ temporal_scale_num: int = 2,
1068
+ slicing_up_num: int = 0,
1069
+ gradient_checkpoint: bool = False,
1070
+ inflation_mode: _inflation_mode_t = "tail",
1071
+ time_receptive_field: _receptive_field_t = "full",
1072
+ slicing_sample_min_size: int = 32,
1073
+ use_quant_conv: bool = True,
1074
+ use_post_quant_conv: bool = True,
1075
+ *args,
1076
+ **kwargs,
1077
+ ):
1078
+ extra_cond_dim = kwargs.pop("extra_cond_dim") if "extra_cond_dim" in kwargs else None
1079
+ self.slicing_sample_min_size = slicing_sample_min_size
1080
+ self.slicing_latent_min_size = slicing_sample_min_size // (2**temporal_scale_num)
1081
+
1082
+ super().__init__(
1083
+ in_channels=in_channels,
1084
+ out_channels=out_channels,
1085
+ # [Override] make sure it can be normally initialized
1086
+ down_block_types=tuple(
1087
+ [down_block_type.replace("3D", "2D") for down_block_type in down_block_types]
1088
+ ),
1089
+ up_block_types=tuple(
1090
+ [up_block_type.replace("3D", "2D") for up_block_type in up_block_types]
1091
+ ),
1092
+ block_out_channels=block_out_channels,
1093
+ layers_per_block=layers_per_block,
1094
+ act_fn=act_fn,
1095
+ latent_channels=latent_channels,
1096
+ norm_num_groups=norm_num_groups,
1097
+ sample_size=sample_size,
1098
+ scaling_factor=scaling_factor,
1099
+ force_upcast=force_upcast,
1100
+ *args,
1101
+ **kwargs,
1102
+ )
1103
+
1104
+ # pass init params to Encoder
1105
+ self.encoder = Encoder3D(
1106
+ in_channels=in_channels,
1107
+ out_channels=latent_channels,
1108
+ down_block_types=down_block_types,
1109
+ block_out_channels=block_out_channels,
1110
+ layers_per_block=layers_per_block,
1111
+ act_fn=act_fn,
1112
+ norm_num_groups=norm_num_groups,
1113
+ double_z=True,
1114
+ extra_cond_dim=extra_cond_dim,
1115
+ # [Override] add temporal_down_num parameter
1116
+ temporal_down_num=temporal_scale_num,
1117
+ gradient_checkpoint=gradient_checkpoint,
1118
+ inflation_mode=inflation_mode,
1119
+ time_receptive_field=time_receptive_field,
1120
+ )
1121
+
1122
+ # pass init params to Decoder
1123
+ self.decoder = Decoder3D(
1124
+ in_channels=latent_channels,
1125
+ out_channels=out_channels,
1126
+ up_block_types=up_block_types,
1127
+ block_out_channels=block_out_channels,
1128
+ layers_per_block=layers_per_block,
1129
+ norm_num_groups=norm_num_groups,
1130
+ act_fn=act_fn,
1131
+ # [Override] add temporal_up_num parameter
1132
+ temporal_up_num=temporal_scale_num,
1133
+ slicing_up_num=slicing_up_num,
1134
+ gradient_checkpoint=gradient_checkpoint,
1135
+ inflation_mode=inflation_mode,
1136
+ time_receptive_field=time_receptive_field,
1137
+ )
1138
+
1139
+ self.quant_conv = (
1140
+ init_causal_conv3d(
1141
+ in_channels=2 * latent_channels,
1142
+ out_channels=2 * latent_channels,
1143
+ kernel_size=1,
1144
+ inflation_mode=inflation_mode,
1145
+ )
1146
+ if use_quant_conv
1147
+ else None
1148
+ )
1149
+ self.post_quant_conv = (
1150
+ init_causal_conv3d(
1151
+ in_channels=latent_channels,
1152
+ out_channels=latent_channels,
1153
+ kernel_size=1,
1154
+ inflation_mode=inflation_mode,
1155
+ )
1156
+ if use_post_quant_conv
1157
+ else None
1158
+ )
1159
+
1160
+ # A hacky way to remove attention.
1161
+ if not attention:
1162
+ self.encoder.mid_block.attentions = torch.nn.ModuleList([None])
1163
+ self.decoder.mid_block.attentions = torch.nn.ModuleList([None])
1164
+
1165
+ @apply_forward_hook
1166
+ def encode(self, x: torch.FloatTensor, return_dict: bool = True) -> AutoencoderKLOutput:
1167
+ h = self.slicing_encode(x)
1168
+ posterior = DiagonalGaussianDistribution(h)
1169
+
1170
+ if not return_dict:
1171
+ return (posterior,)
1172
+
1173
+ return AutoencoderKLOutput(latent_dist=posterior)
1174
+
1175
+ @apply_forward_hook
1176
+ def decode(
1177
+ self, z: torch.Tensor, return_dict: bool = True
1178
+ ) -> Union[DecoderOutput, torch.Tensor]:
1179
+ decoded = self.slicing_decode(z)
1180
+
1181
+ if not return_dict:
1182
+ return (decoded,)
1183
+
1184
+ return DecoderOutput(sample=decoded)
1185
+
1186
+ def _encode(
1187
+ self, x: torch.Tensor, memory_state: MemoryState = MemoryState.DISABLED
1188
+ ) -> torch.Tensor:
1189
+ _x = x.to(self.device)
1190
+ _x = causal_conv_slice_inputs(_x, self.slicing_sample_min_size, memory_state=memory_state)
1191
+ h = self.encoder(_x, memory_state=memory_state)
1192
+ if self.quant_conv is not None:
1193
+ output = self.quant_conv(h, memory_state=memory_state)
1194
+ else:
1195
+ output = h
1196
+ output = causal_conv_gather_outputs(output)
1197
+ return output.to(x.device)
1198
+
1199
+ def _decode(
1200
+ self, z: torch.Tensor, memory_state: MemoryState = MemoryState.DISABLED
1201
+ ) -> torch.Tensor:
1202
+ _z = z.to(self.device)
1203
+ _z = causal_conv_slice_inputs(_z, self.slicing_latent_min_size, memory_state=memory_state)
1204
+ if self.post_quant_conv is not None:
1205
+ _z = self.post_quant_conv(_z, memory_state=memory_state)
1206
+ output = self.decoder(_z, memory_state=memory_state)
1207
+ output = causal_conv_gather_outputs(output)
1208
+ return output.to(z.device)
1209
+
1210
+ def slicing_encode(self, x: torch.Tensor) -> torch.Tensor:
1211
+ sp_size = get_sequence_parallel_world_size()
1212
+ if self.use_slicing and (x.shape[2] - 1) > self.slicing_sample_min_size * sp_size:
1213
+ x_slices = x[:, :, 1:].split(split_size=self.slicing_sample_min_size * sp_size, dim=2)
1214
+ encoded_slices = [
1215
+ self._encode(
1216
+ torch.cat((x[:, :, :1], x_slices[0]), dim=2),
1217
+ memory_state=MemoryState.INITIALIZING,
1218
+ )
1219
+ ]
1220
+ for x_idx in range(1, len(x_slices)):
1221
+ encoded_slices.append(
1222
+ self._encode(x_slices[x_idx], memory_state=MemoryState.ACTIVE)
1223
+ )
1224
+ return torch.cat(encoded_slices, dim=2)
1225
+ else:
1226
+ return self._encode(x)
1227
+
1228
+ def slicing_decode(self, z: torch.Tensor) -> torch.Tensor:
1229
+ sp_size = get_sequence_parallel_world_size()
1230
+ if self.use_slicing and (z.shape[2] - 1) > self.slicing_latent_min_size * sp_size:
1231
+ z_slices = z[:, :, 1:].split(split_size=self.slicing_latent_min_size * sp_size, dim=2)
1232
+ decoded_slices = [
1233
+ self._decode(
1234
+ torch.cat((z[:, :, :1], z_slices[0]), dim=2),
1235
+ memory_state=MemoryState.INITIALIZING,
1236
+ )
1237
+ ]
1238
+ for z_idx in range(1, len(z_slices)):
1239
+ decoded_slices.append(
1240
+ self._decode(z_slices[z_idx], memory_state=MemoryState.ACTIVE)
1241
+ )
1242
+ return torch.cat(decoded_slices, dim=2)
1243
+ else:
1244
+ return self._decode(z)
1245
+
1246
+ def tiled_encode(self, x: torch.Tensor, **kwargs) -> torch.Tensor:
1247
+ raise NotImplementedError
1248
+
1249
+ def tiled_decode(self, z: torch.Tensor, **kwargs) -> torch.Tensor:
1250
+ raise NotImplementedError
1251
+
1252
+ def forward(
1253
+ self, x: torch.FloatTensor, mode: Literal["encode", "decode", "all"] = "all", **kwargs
1254
+ ):
1255
+ # x: [b c t h w]
1256
+ if mode == "encode":
1257
+ h = self.encode(x)
1258
+ return h.latent_dist
1259
+ elif mode == "decode":
1260
+ h = self.decode(x)
1261
+ return h.sample
1262
+ else:
1263
+ h = self.encode(x)
1264
+ h = self.decode(h.latent_dist.mode())
1265
+ return h.sample
1266
+
1267
+ def load_state_dict(self, state_dict, strict=False):
1268
+ # Newer version of diffusers changed the model keys,
1269
+ # causing incompatibility with old checkpoints.
1270
+ # They provided a method for conversion.
1271
+ # We call conversion before loading state_dict.
1272
+ convert_deprecated_attention_blocks = getattr(
1273
+ self, "_convert_deprecated_attention_blocks", None
1274
+ )
1275
+ if callable(convert_deprecated_attention_blocks):
1276
+ convert_deprecated_attention_blocks(state_dict)
1277
+ return super().load_state_dict(state_dict, strict)
1278
+
1279
+
1280
+ class VideoAutoencoderKLWrapper(VideoAutoencoderKL):
1281
+ def __init__(
1282
+ self,
1283
+ *args,
1284
+ spatial_downsample_factor: int,
1285
+ temporal_downsample_factor: int,
1286
+ freeze_encoder: bool,
1287
+ **kwargs,
1288
+ ):
1289
+ self.spatial_downsample_factor = spatial_downsample_factor
1290
+ self.temporal_downsample_factor = temporal_downsample_factor
1291
+ self.freeze_encoder = freeze_encoder
1292
+ super().__init__(*args, **kwargs)
1293
+
1294
+ def forward(self, x: torch.FloatTensor) -> CausalAutoencoderOutput:
1295
+ with torch.no_grad() if self.freeze_encoder else nullcontext():
1296
+ z, p = self.encode(x)
1297
+ x = self.decode(z).sample
1298
+ return CausalAutoencoderOutput(x, z, p)
1299
+
1300
+ def encode(self, x: torch.FloatTensor) -> CausalEncoderOutput:
1301
+ if x.ndim == 4:
1302
+ x = x.unsqueeze(2)
1303
+ p = super().encode(x).latent_dist
1304
+ z = p.sample().squeeze(2)
1305
+ return CausalEncoderOutput(z, p)
1306
+
1307
+ def decode(self, z: torch.FloatTensor) -> CausalDecoderOutput:
1308
+ if z.ndim == 4:
1309
+ z = z.unsqueeze(2)
1310
+ x = super().decode(z).sample.squeeze(2)
1311
+ return CausalDecoderOutput(x)
1312
+
1313
+ def preprocess(self, x: torch.Tensor):
1314
+ # x should in [B, C, T, H, W], [B, C, H, W]
1315
+ assert x.ndim == 4 or x.size(2) % 4 == 1
1316
+ return x
1317
+
1318
+ def postprocess(self, x: torch.Tensor):
1319
+ # x should in [B, C, T, H, W], [B, C, H, W]
1320
+ return x
1321
+
1322
+ def set_causal_slicing(
1323
+ self,
1324
+ *,
1325
+ split_size: Optional[int],
1326
+ memory_device: _memory_device_t,
1327
+ ):
1328
+ assert (
1329
+ split_size is None or memory_device is not None
1330
+ ), "if split_size is set, memory_device must not be None."
1331
+ if split_size is not None:
1332
+ self.enable_slicing()
1333
+ self.slicing_sample_min_size = split_size
1334
+ self.slicing_latent_min_size = split_size // self.temporal_downsample_factor
1335
+ else:
1336
+ self.disable_slicing()
1337
+ for module in self.modules():
1338
+ if isinstance(module, InflatedCausalConv3d):
1339
+ module.set_memory_device(memory_device)
1340
+
1341
+ def set_memory_limit(self, conv_max_mem: Optional[float], norm_max_mem: Optional[float]):
1342
+ set_norm_limit(norm_max_mem)
1343
+ for m in self.modules():
1344
+ if isinstance(m, InflatedCausalConv3d):
1345
+ m.set_memory_limit(conv_max_mem if conv_max_mem is not None else float("inf"))
models/video_vae_v3/modules/causal_inflation_lib.py ADDED
@@ -0,0 +1,460 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # // Copyright (c) 2025 Bytedance Ltd. and/or its affiliates
2
+ # //
3
+ # // Licensed under the Apache License, Version 2.0 (the "License");
4
+ # // you may not use this file except in compliance with the License.
5
+ # // You may obtain a copy of the License at
6
+ # //
7
+ # // http://www.apache.org/licenses/LICENSE-2.0
8
+ # //
9
+ # // Unless required by applicable law or agreed to in writing, software
10
+ # // distributed under the License is distributed on an "AS IS" BASIS,
11
+ # // WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # // See the License for the specific language governing permissions and
13
+ # // limitations under the License.
14
+
15
+ import math
16
+ from contextlib import contextmanager
17
+ from typing import List, Optional, Union
18
+ import torch
19
+ import torch.distributed as dist
20
+ import torch.nn.functional as F
21
+ from diffusers.models.normalization import RMSNorm
22
+ from einops import rearrange
23
+ from torch import Tensor, nn
24
+ from torch.nn import Conv3d
25
+
26
+ from common.distributed.advanced import (
27
+ get_next_sequence_parallel_rank,
28
+ get_prev_sequence_parallel_rank,
29
+ get_sequence_parallel_group,
30
+ get_sequence_parallel_rank,
31
+ get_sequence_parallel_world_size,
32
+ )
33
+ from common.logger import get_logger
34
+ from models.video_vae_v3.modules.context_parallel_lib import cache_send_recv, get_cache_size
35
+ from models.video_vae_v3.modules.global_config import get_norm_limit
36
+ from models.video_vae_v3.modules.types import MemoryState, _inflation_mode_t, _memory_device_t
37
+
38
+ logger = get_logger(__name__)
39
+
40
+
41
+ @contextmanager
42
+ def ignore_padding(model):
43
+ orig_padding = model.padding
44
+ model.padding = (0, 0, 0)
45
+ try:
46
+ yield
47
+ finally:
48
+ model.padding = orig_padding
49
+
50
+
51
+ class InflatedCausalConv3d(Conv3d):
52
+ def __init__(
53
+ self,
54
+ *args,
55
+ inflation_mode: _inflation_mode_t,
56
+ memory_device: _memory_device_t = "same",
57
+ **kwargs,
58
+ ):
59
+ self.inflation_mode = inflation_mode
60
+ self.memory = None
61
+ super().__init__(*args, **kwargs)
62
+ self.temporal_padding = self.padding[0]
63
+ self.memory_device = memory_device
64
+ self.padding = (0, *self.padding[1:]) # Remove temporal pad to keep causal.
65
+ self.memory_limit = float("inf")
66
+
67
+ def set_memory_limit(self, value: float):
68
+ self.memory_limit = value
69
+
70
+ def set_memory_device(self, memory_device: _memory_device_t):
71
+ self.memory_device = memory_device
72
+
73
+ def memory_limit_conv(
74
+ self,
75
+ x,
76
+ *,
77
+ split_dim=3,
78
+ padding=(0, 0, 0, 0, 0, 0),
79
+ prev_cache=None,
80
+ ):
81
+ # Compatible with no limit.
82
+ if math.isinf(self.memory_limit):
83
+ if prev_cache is not None:
84
+ x = torch.cat([prev_cache, x], dim=split_dim - 1)
85
+ return super().forward(x)
86
+
87
+ # Compute tensor shape after concat & padding.
88
+ shape = torch.tensor(x.size())
89
+ if prev_cache is not None:
90
+ shape[split_dim - 1] += prev_cache.size(split_dim - 1)
91
+ shape[-3:] += torch.tensor(padding).view(3, 2).sum(-1).flip(0)
92
+ memory_occupy = shape.prod() * x.element_size() / 1024**3 # GiB
93
+ logger.debug(
94
+ f"x:{(shape, x.dtype)} {memory_occupy:.3f}GiB "
95
+ f"prev_cache:{prev_cache.shape if prev_cache is not None else None}"
96
+ )
97
+ if memory_occupy < self.memory_limit or split_dim == x.ndim:
98
+ if prev_cache is not None:
99
+ x = torch.cat([prev_cache, x], dim=split_dim - 1)
100
+ x = F.pad(x, padding, value=0.0)
101
+ with ignore_padding(self):
102
+ return super().forward(x)
103
+
104
+ logger.debug(
105
+ f"Exceed memory limit {memory_occupy} > {self.memory_limit}, split dim {split_dim}"
106
+ )
107
+
108
+ # Split input (& prev_cache).
109
+ num_splits = math.ceil(memory_occupy / self.memory_limit)
110
+ size_per_split = x.size(split_dim) // num_splits
111
+ split_sizes = [size_per_split] * (num_splits - 1)
112
+ split_sizes += [x.size(split_dim) - sum(split_sizes)]
113
+
114
+ x = list(x.split(split_sizes, dim=split_dim))
115
+ logger.debug(f"Conv inputs: {[inp.size() for inp in x]} {x[0].dtype}")
116
+ if prev_cache is not None:
117
+ prev_cache = list(prev_cache.split(split_sizes, dim=split_dim))
118
+
119
+ # Loop Fwd.
120
+ cache = None
121
+ for idx in range(len(x)):
122
+ # Concat prev cache from last dim
123
+ if prev_cache is not None:
124
+ x[idx] = torch.cat([prev_cache[idx], x[idx]], dim=split_dim - 1)
125
+
126
+ # Get padding pattern.
127
+ lpad_dim = (x[idx].ndim - split_dim - 1) * 2
128
+ rpad_dim = lpad_dim + 1
129
+ padding = list(padding)
130
+ padding[lpad_dim] = self.padding[split_dim - 2] if idx == 0 else 0
131
+ padding[rpad_dim] = self.padding[split_dim - 2] if idx == len(x) - 1 else 0
132
+ pad_len = padding[lpad_dim] + padding[rpad_dim]
133
+ padding = tuple(padding)
134
+
135
+ # Prepare cache for next slice (this dim).
136
+ next_cache = None
137
+ cache_len = cache.size(split_dim) if cache is not None else 0
138
+ next_catch_size = get_cache_size(
139
+ conv_module=self,
140
+ input_len=x[idx].size(split_dim) + cache_len,
141
+ pad_len=pad_len,
142
+ dim=split_dim - 2,
143
+ )
144
+ if next_catch_size != 0:
145
+ assert next_catch_size <= x[idx].size(split_dim)
146
+ next_cache = (
147
+ x[idx].transpose(0, split_dim)[-next_catch_size:].transpose(0, split_dim)
148
+ )
149
+
150
+ # Recursive.
151
+ x[idx] = self.memory_limit_conv(
152
+ x[idx],
153
+ split_dim=split_dim + 1,
154
+ padding=padding,
155
+ prev_cache=cache,
156
+ )
157
+
158
+ # Update cache.
159
+ cache = next_cache
160
+
161
+ logger.debug(f"Conv outputs, concat(dim={split_dim}): {[d.size() for d in x]}")
162
+ return torch.cat(x, split_dim)
163
+
164
+ def forward(
165
+ self,
166
+ input: Union[Tensor, List[Tensor]],
167
+ memory_state: MemoryState = MemoryState.UNSET,
168
+ ) -> Tensor:
169
+ assert memory_state != MemoryState.UNSET
170
+ if memory_state != MemoryState.ACTIVE:
171
+ self.memory = None
172
+ if (
173
+ math.isinf(self.memory_limit)
174
+ and torch.is_tensor(input)
175
+ and get_sequence_parallel_group() is None
176
+ ):
177
+ return self.basic_forward(input, memory_state)
178
+ return self.slicing_forward(input, memory_state)
179
+
180
+ def basic_forward(self, input: Tensor, memory_state: MemoryState = MemoryState.UNSET):
181
+ mem_size = self.stride[0] - self.kernel_size[0]
182
+ if (self.memory is not None) and (memory_state == MemoryState.ACTIVE):
183
+ input = extend_head(input, memory=self.memory, times=-1)
184
+ else:
185
+ input = extend_head(input, times=self.temporal_padding * 2)
186
+ memory = (
187
+ input[:, :, mem_size:].detach()
188
+ if (mem_size != 0 and memory_state != MemoryState.DISABLED)
189
+ else None
190
+ )
191
+ if (
192
+ memory_state != MemoryState.DISABLED
193
+ and not self.training
194
+ and (self.memory_device is not None)
195
+ ):
196
+ self.memory = memory
197
+ if self.memory_device == "cpu" and self.memory is not None:
198
+ self.memory = self.memory.to("cpu")
199
+ return super().forward(input)
200
+
201
+ def slicing_forward(
202
+ self,
203
+ input: Union[Tensor, List[Tensor]],
204
+ memory_state: MemoryState = MemoryState.UNSET,
205
+ ) -> Tensor:
206
+ squeeze_out = False
207
+ if torch.is_tensor(input):
208
+ input = [input]
209
+ squeeze_out = True
210
+
211
+ cache_size = self.kernel_size[0] - self.stride[0]
212
+ cache = cache_send_recv(
213
+ input, cache_size=cache_size, memory=self.memory, times=self.temporal_padding * 2
214
+ )
215
+
216
+ # For slice=4 and sp=2, and 17 frames in total
217
+ # sp0 sp1
218
+ # slice 0: [`0 0` 0 1 2 {3 4}] [{3 4} 5 6 (7 8)] extend=`0 0` cache={3 4} memory=(7 8)
219
+ # slice 1: [(7 8) 9 10 {11 12}] [{11 12} 13 14 15 16]
220
+ sp_rank = get_sequence_parallel_rank()
221
+ sp_size = get_sequence_parallel_world_size()
222
+ sp_group = get_sequence_parallel_group()
223
+ send_dst = get_next_sequence_parallel_rank()
224
+ recv_src = get_prev_sequence_parallel_rank()
225
+ if (
226
+ memory_state in [MemoryState.INITIALIZING, MemoryState.ACTIVE] # use_slicing
227
+ and not self.training
228
+ and (self.memory_device is not None)
229
+ and sp_rank in [0, sp_size - 1]
230
+ and cache_size != 0
231
+ ):
232
+ if cache_size > input[-1].size(2) and cache is not None and len(input) == 1:
233
+ input[0] = torch.cat([cache, input[0]], dim=2)
234
+ cache = None
235
+ assert cache_size <= input[-1].size(2)
236
+ if sp_size == 1:
237
+ self.memory = input[-1][:, :, -cache_size:].detach().contiguous()
238
+ else:
239
+ if sp_rank == sp_size - 1:
240
+ dist.send(
241
+ input[-1][:, :, -cache_size:].detach().contiguous(),
242
+ send_dst,
243
+ group=sp_group,
244
+ )
245
+ if sp_rank == 0:
246
+ shape = list(input[0].size())
247
+ shape[2] = cache_size
248
+ self.memory = torch.empty(
249
+ *shape, device=input[0].device, dtype=input[0].dtype
250
+ ).contiguous()
251
+ dist.recv(self.memory, recv_src, group=sp_group)
252
+ if self.memory_device == "cpu" and self.memory is not None:
253
+ self.memory = self.memory.to("cpu")
254
+
255
+ padding = tuple(x for x in reversed(self.padding) for _ in range(2))
256
+ for i in range(len(input)):
257
+ # Prepare cache for next input slice.
258
+ next_cache = None
259
+ cache_size = 0
260
+ if i < len(input) - 1:
261
+ cache_len = cache.size(2) if cache is not None else 0
262
+ cache_size = get_cache_size(self, input[i].size(2) + cache_len, pad_len=0)
263
+ if cache_size != 0:
264
+ if cache_size > input[i].size(2) and cache is not None:
265
+ input[i] = torch.cat([cache, input[i]], dim=2)
266
+ cache = None
267
+ assert cache_size <= input[i].size(2), f"{cache_size} > {input[i].size(2)}"
268
+ next_cache = input[i][:, :, -cache_size:]
269
+
270
+ # Conv forward for this input slice.
271
+ input[i] = self.memory_limit_conv(
272
+ input[i],
273
+ padding=padding,
274
+ prev_cache=cache,
275
+ )
276
+
277
+ # Update cache.
278
+ cache = next_cache
279
+
280
+ return input[0] if squeeze_out else input
281
+
282
+ def tflops(self, args, kwargs, output) -> float:
283
+ if torch.is_tensor(output):
284
+ output_numel = output.numel()
285
+ elif isinstance(output, list):
286
+ output_numel = sum(o.numel() for o in output)
287
+ else:
288
+ raise NotImplementedError
289
+ return (2 * math.prod(self.kernel_size) * self.in_channels * (output_numel / 1e6)) / 1e6
290
+
291
+ def _load_from_state_dict(
292
+ self, state_dict, prefix, local_metadata, strict, missing_keys, unexpected_keys, error_msgs
293
+ ):
294
+ if self.inflation_mode != "none":
295
+ state_dict = modify_state_dict(
296
+ self,
297
+ state_dict,
298
+ prefix,
299
+ inflate_weight_fn=inflate_weight,
300
+ inflate_bias_fn=inflate_bias,
301
+ )
302
+ super()._load_from_state_dict(
303
+ state_dict,
304
+ prefix,
305
+ local_metadata,
306
+ (strict and self.inflation_mode == "none"),
307
+ missing_keys,
308
+ unexpected_keys,
309
+ error_msgs,
310
+ )
311
+
312
+
313
+ def init_causal_conv3d(
314
+ *args,
315
+ inflation_mode: _inflation_mode_t,
316
+ **kwargs,
317
+ ):
318
+ """
319
+ Initialize a Causal-3D convolution layer.
320
+ Parameters:
321
+ inflation_mode: Listed as below. It's compatible with all the 3D-VAE checkpoints we have.
322
+ - none: No inflation will be conducted.
323
+ The loading logic of state dict will fall back to default.
324
+ - tail / replicate: Refer to the definition of `InflatedCausalConv3d`.
325
+ """
326
+ return InflatedCausalConv3d(*args, inflation_mode=inflation_mode, **kwargs)
327
+
328
+
329
+ def causal_norm_wrapper(norm_layer: nn.Module, x: torch.Tensor) -> torch.Tensor:
330
+ input_dtype = x.dtype
331
+ if isinstance(norm_layer, (nn.LayerNorm, RMSNorm)):
332
+ if x.ndim == 4:
333
+ x = rearrange(x, "b c h w -> b h w c")
334
+ x = norm_layer(x)
335
+ x = rearrange(x, "b h w c -> b c h w")
336
+ return x.to(input_dtype)
337
+ if x.ndim == 5:
338
+ x = rearrange(x, "b c t h w -> b t h w c")
339
+ x = norm_layer(x)
340
+ x = rearrange(x, "b t h w c -> b c t h w")
341
+ return x.to(input_dtype)
342
+ if isinstance(norm_layer, (nn.GroupNorm, nn.BatchNorm2d, nn.SyncBatchNorm)):
343
+ if x.ndim <= 4:
344
+ return norm_layer(x).to(input_dtype)
345
+ if x.ndim == 5:
346
+ t = x.size(2)
347
+ x = rearrange(x, "b c t h w -> (b t) c h w")
348
+ memory_occupy = x.numel() * x.element_size() / 1024**3
349
+ if isinstance(norm_layer, nn.GroupNorm) and memory_occupy > get_norm_limit():
350
+ num_chunks = min(4 if x.element_size() == 2 else 2, norm_layer.num_groups)
351
+ logger.debug(f"large tensor {x.shape}, norm in {num_chunks} chunks")
352
+ assert norm_layer.num_groups % num_chunks == 0
353
+ num_groups_per_chunk = norm_layer.num_groups // num_chunks
354
+
355
+ x = list(x.chunk(num_chunks, dim=1))
356
+ weights = norm_layer.weight.chunk(num_chunks, dim=0)
357
+ biases = norm_layer.bias.chunk(num_chunks, dim=0)
358
+ for i, (w, b) in enumerate(zip(weights, biases)):
359
+ x[i] = F.group_norm(x[i], num_groups_per_chunk, w, b, norm_layer.eps)
360
+ x[i] = x[i].to(input_dtype)
361
+ x = torch.cat(x, dim=1)
362
+ else:
363
+ x = norm_layer(x)
364
+ x = rearrange(x, "(b t) c h w -> b c t h w", t=t)
365
+ return x.to(input_dtype)
366
+ raise NotImplementedError
367
+
368
+
369
+ def remove_head(tensor: Tensor, times: int = 1) -> Tensor:
370
+ """
371
+ Remove duplicated first frame features in the up-sampling process.
372
+ """
373
+ sp_rank = get_sequence_parallel_rank()
374
+ if times == 0 or sp_rank > 0:
375
+ return tensor
376
+ return torch.cat(tensors=(tensor[:, :, :1], tensor[:, :, times + 1 :]), dim=2)
377
+
378
+
379
+ def extend_head(tensor: Tensor, times: int = 2, memory: Optional[Tensor] = None) -> Tensor:
380
+ """
381
+ When memory is None:
382
+ - Duplicate first frame features in the down-sampling process.
383
+ When memory is not None:
384
+ - Concatenate memory features with the input features to keep temporal consistency.
385
+ """
386
+ if memory is not None:
387
+ return torch.cat((memory.to(tensor), tensor), dim=2)
388
+ assert times >= 0, "Invalid input for function 'extend_head'!"
389
+ if times == 0:
390
+ return tensor
391
+ else:
392
+ tile_repeat = [1] * tensor.ndim
393
+ tile_repeat[2] = times
394
+ return torch.cat(tensors=(torch.tile(tensor[:, :, :1], tile_repeat), tensor), dim=2)
395
+
396
+
397
+ def inflate_weight(weight_2d: torch.Tensor, weight_3d: torch.Tensor, inflation_mode: str):
398
+ """
399
+ Inflate a 2D convolution weight matrix to a 3D one.
400
+ Parameters:
401
+ weight_2d: The weight matrix of 2D conv to be inflated.
402
+ weight_3d: The weight matrix of 3D conv to be initialized.
403
+ inflation_mode: the mode of inflation
404
+ """
405
+ assert inflation_mode in ["tail", "replicate"]
406
+ assert weight_3d.shape[:2] == weight_2d.shape[:2]
407
+ with torch.no_grad():
408
+ if inflation_mode == "replicate":
409
+ depth = weight_3d.size(2)
410
+ weight_3d.copy_(weight_2d.unsqueeze(2).repeat(1, 1, depth, 1, 1) / depth)
411
+ else:
412
+ weight_3d.fill_(0.0)
413
+ weight_3d[:, :, -1].copy_(weight_2d)
414
+ return weight_3d
415
+
416
+
417
+ def inflate_bias(bias_2d: torch.Tensor, bias_3d: torch.Tensor, inflation_mode: str):
418
+ """
419
+ Inflate a 2D convolution bias tensor to a 3D one
420
+ Parameters:
421
+ bias_2d: The bias tensor of 2D conv to be inflated.
422
+ bias_3d: The bias tensor of 3D conv to be initialized.
423
+ inflation_mode: Placeholder to align `inflate_weight`.
424
+ """
425
+ assert bias_3d.shape == bias_2d.shape
426
+ with torch.no_grad():
427
+ bias_3d.copy_(bias_2d)
428
+ return bias_3d
429
+
430
+
431
+ def modify_state_dict(layer, state_dict, prefix, inflate_weight_fn, inflate_bias_fn):
432
+ """
433
+ the main function to inflated 2D parameters to 3D.
434
+ """
435
+ weight_name = prefix + "weight"
436
+ bias_name = prefix + "bias"
437
+ if weight_name in state_dict:
438
+ weight_2d = state_dict[weight_name]
439
+ if weight_2d.dim() == 4:
440
+ # Assuming the 2D weights are 4D tensors (out_channels, in_channels, h, w)
441
+ weight_3d = inflate_weight_fn(
442
+ weight_2d=weight_2d,
443
+ weight_3d=layer.weight,
444
+ inflation_mode=layer.inflation_mode,
445
+ )
446
+ state_dict[weight_name] = weight_3d
447
+ else:
448
+ return state_dict
449
+ # It's a 3d state dict, should not do inflation on both bias and weight.
450
+ if bias_name in state_dict:
451
+ bias_2d = state_dict[bias_name]
452
+ if bias_2d.dim() == 1:
453
+ # Assuming the 2D biases are 1D tensors (out_channels,)
454
+ bias_3d = inflate_bias_fn(
455
+ bias_2d=bias_2d,
456
+ bias_3d=layer.bias,
457
+ inflation_mode=layer.inflation_mode,
458
+ )
459
+ state_dict[bias_name] = bias_3d
460
+ return state_dict
models/video_vae_v3/modules/context_parallel_lib.py ADDED
@@ -0,0 +1,164 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # // Copyright (c) 2025 Bytedance Ltd. and/or its affiliates
2
+ # //
3
+ # // Licensed under the Apache License, Version 2.0 (the "License");
4
+ # // you may not use this file except in compliance with the License.
5
+ # // You may obtain a copy of the License at
6
+ # //
7
+ # // http://www.apache.org/licenses/LICENSE-2.0
8
+ # //
9
+ # // Unless required by applicable law or agreed to in writing, software
10
+ # // distributed under the License is distributed on an "AS IS" BASIS,
11
+ # // WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # // See the License for the specific language governing permissions and
13
+ # // limitations under the License.
14
+
15
+ from typing import List
16
+ import torch
17
+ import torch.distributed as dist
18
+ import torch.nn.functional as F
19
+ from torch import Tensor
20
+
21
+ from common.distributed import get_device
22
+ from common.distributed.advanced import (
23
+ get_next_sequence_parallel_rank,
24
+ get_prev_sequence_parallel_rank,
25
+ get_sequence_parallel_group,
26
+ get_sequence_parallel_rank,
27
+ get_sequence_parallel_world_size,
28
+ )
29
+ from common.distributed.ops import Gather
30
+ from common.logger import get_logger
31
+ from models.video_vae_v3.modules.types import MemoryState
32
+
33
+ logger = get_logger(__name__)
34
+
35
+
36
+ def causal_conv_slice_inputs(x, split_size, memory_state):
37
+ sp_size = get_sequence_parallel_world_size()
38
+ sp_group = get_sequence_parallel_group()
39
+ sp_rank = get_sequence_parallel_rank()
40
+ if sp_group is None:
41
+ return x
42
+
43
+ assert memory_state != MemoryState.UNSET
44
+ leave_out = 1 if memory_state != MemoryState.ACTIVE else 0
45
+
46
+ # Should have at least sp_size slices.
47
+ num_slices = (x.size(2) - leave_out) // split_size
48
+ assert num_slices >= sp_size, f"{num_slices} < {sp_size}"
49
+
50
+ split_sizes = [split_size + leave_out] + [split_size] * (num_slices - 1)
51
+ split_sizes += [x.size(2) - sum(split_sizes)]
52
+ assert sum(split_sizes) == x.size(2)
53
+
54
+ split_sizes = torch.tensor(split_sizes)
55
+ slices_per_rank = len(split_sizes) // sp_size
56
+ split_sizes = split_sizes.split(
57
+ [slices_per_rank] * (sp_size - 1) + [len(split_sizes) - slices_per_rank * (sp_size - 1)]
58
+ )
59
+ split_sizes = list(map(lambda s: s.sum().item(), split_sizes))
60
+ logger.debug(f"split_sizes: {split_sizes}")
61
+ return x.split(split_sizes, dim=2)[sp_rank]
62
+
63
+
64
+ def causal_conv_gather_outputs(x):
65
+ sp_group = get_sequence_parallel_group()
66
+ sp_size = get_sequence_parallel_world_size()
67
+ if sp_group is None:
68
+ return x
69
+
70
+ # Communicate shapes.
71
+ unpad_lens = torch.empty((sp_size,), device=get_device(), dtype=torch.long)
72
+ local_unpad_len = torch.tensor([x.size(2)], device=get_device(), dtype=torch.long)
73
+ torch.distributed.all_gather_into_tensor(unpad_lens, local_unpad_len, group=sp_group)
74
+
75
+ # Padding to max_len for gather.
76
+ max_len = unpad_lens.max()
77
+ x_pad = F.pad(x, (0, 0, 0, 0, 0, max_len - x.size(2))).contiguous()
78
+
79
+ # Gather outputs.
80
+ x_pad = Gather.apply(sp_group, x_pad, 2, True)
81
+
82
+ # Remove padding.
83
+ x_pad_lists = list(x_pad.chunk(sp_size, dim=2))
84
+ for i, (x_pad, unpad_len) in enumerate(zip(x_pad_lists, unpad_lens)):
85
+ x_pad_lists[i] = x_pad[:, :, :unpad_len]
86
+
87
+ return torch.cat(x_pad_lists, dim=2)
88
+
89
+
90
+ def get_output_len(conv_module, input_len, pad_len, dim=0):
91
+ dilated_kernerl_size = conv_module.dilation[dim] * (conv_module.kernel_size[dim] - 1) + 1
92
+ output_len = (input_len + pad_len - dilated_kernerl_size) // conv_module.stride[dim] + 1
93
+ return output_len
94
+
95
+
96
+ def get_cache_size(conv_module, input_len, pad_len, dim=0):
97
+ dilated_kernerl_size = conv_module.dilation[dim] * (conv_module.kernel_size[dim] - 1) + 1
98
+ output_len = (input_len + pad_len - dilated_kernerl_size) // conv_module.stride[dim] + 1
99
+ remain_len = (
100
+ input_len + pad_len - ((output_len - 1) * conv_module.stride[dim] + dilated_kernerl_size)
101
+ )
102
+ overlap_len = dilated_kernerl_size - conv_module.stride[dim]
103
+ cache_len = overlap_len + remain_len # >= 0
104
+ logger.debug(
105
+ f"I:{input_len}, "
106
+ f"P:{pad_len}, "
107
+ f"K:{conv_module.kernel_size[dim]}, "
108
+ f"S:{conv_module.stride[dim]}, "
109
+ f"O:{output_len}, "
110
+ f"Cache:{cache_len}"
111
+ )
112
+ assert output_len > 0
113
+ return cache_len
114
+
115
+
116
+ def cache_send_recv(tensor: List[Tensor], cache_size, times, memory=None):
117
+ sp_group = get_sequence_parallel_group()
118
+ sp_rank = get_sequence_parallel_rank()
119
+ sp_size = get_sequence_parallel_world_size()
120
+ send_dst = get_next_sequence_parallel_rank()
121
+ recv_src = get_prev_sequence_parallel_rank()
122
+ recv_buffer = None
123
+ recv_req = None
124
+
125
+ logger.debug(
126
+ f"[sp{sp_rank}] cur_tensors:{[(t.size(), t.dtype) for t in tensor]}, times: {times}"
127
+ )
128
+ if sp_rank == 0 or sp_group is None:
129
+ if memory is not None:
130
+ recv_buffer = memory.to(tensor[0])
131
+ elif times > 0:
132
+ tile_repeat = [1] * tensor[0].ndim
133
+ tile_repeat[2] = times
134
+ recv_buffer = torch.tile(tensor[0][:, :, :1], tile_repeat)
135
+
136
+ if cache_size != 0 and sp_group is not None:
137
+ if sp_rank > 0:
138
+ shape = list(tensor[0].size())
139
+ shape[2] = cache_size
140
+ recv_buffer = torch.empty(
141
+ *shape, device=tensor[0].device, dtype=tensor[0].dtype
142
+ ).contiguous()
143
+ recv_req = dist.irecv(recv_buffer, recv_src, group=sp_group)
144
+ if sp_rank < sp_size - 1:
145
+ if cache_size > tensor[-1].size(2) and len(tensor) == 1:
146
+ logger.debug(f"[sp{sp_rank}] force concat before send {tensor[-1].size()}")
147
+ if recv_req is not None:
148
+ recv_req.wait()
149
+ tensor[0] = torch.cat([recv_buffer, tensor[0]], dim=2)
150
+ recv_buffer = None
151
+ assert cache_size <= tensor[-1].size(
152
+ 2
153
+ ), f"Not enough value to cache, got {tensor[-1].size()}, cache_size={cache_size}"
154
+ dist.isend(
155
+ tensor[-1][:, :, -cache_size:].detach().contiguous(), send_dst, group=sp_group
156
+ )
157
+ if recv_req is not None:
158
+ recv_req.wait()
159
+
160
+ logger.debug(
161
+ f"[sp{sp_rank}] recv_src:{recv_src}, "
162
+ f"recv_buffer:{recv_buffer.size() if recv_buffer is not None else None}"
163
+ )
164
+ return recv_buffer