Rbaerk commited on
Commit
71dfe4d
·
verified ·
1 Parent(s): 5c39aa1

Release TiTok motion encoder code and checkpoint

Browse files
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 [2023] [Zhongjie Duan]
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.
NOTICE ADDED
@@ -0,0 +1,6 @@
 
 
 
 
 
 
 
1
+ IM-Animation motion encoder uses modified TiTok components from
2
+ https://github.com/bytedance/1d-tokenizer (Apache-2.0).
3
+ Original copyright and reference notices are retained in source files.
4
+ This distribution extracts the encoder, vector quantizer, and preprocessing
5
+ from the IM-Animation research implementation and adds an encoder-only loader.
6
+ The original TiTok legacy reshape is preserved.
README.md ADDED
@@ -0,0 +1,87 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ language:
3
+ - en
4
+ license: apache-2.0
5
+ tags:
6
+ - im-animation
7
+ - titok
8
+ - motion-encoder
9
+ - image-feature-extraction
10
+ - arxiv:2602.07498
11
+ ---
12
+
13
+ # IM-Animation
14
+
15
+ [Paper](https://arxiv.org/abs/2602.07498) · [Project page](https://rabberk.github.io/IM-Animation/) · [Motion encoder weights](https://huggingface.co/Rbaerk/IM-Animation-Motion-Encoder)
16
+
17
+ **IM-Animation: An Implicit Motion Representation for Identity-decoupled Character Animation**
18
+
19
+ This release contains the final locally retained **TiTok-based motion encoder** implementation and an exported checkpoint. It provides frame-level motion tokens; it does not include the full animation generator or retargeting network. Code and project videos are available in the linked GitHub repository.
20
+
21
+ ## Installation
22
+
23
+ ```bash
24
+ git clone https://github.com/rabberk/IM-Animation.git
25
+ cd IM-Animation
26
+ pip install -r requirements.txt
27
+ ```
28
+
29
+ The encoder was verified with PyTorch 2.7.1 on CPU using the exported BF16 weights.
30
+
31
+ ## Download weights
32
+
33
+ ```python
34
+ from huggingface_hub import hf_hub_download
35
+
36
+ hf_hub_download(
37
+ repo_id="Rbaerk/IM-Animation-Motion-Encoder",
38
+ filename="motion_encoder_latest.safetensors",
39
+ local_dir=".",
40
+ )
41
+ ```
42
+
43
+ The weight file is 607,017,208 bytes (approximately 579 MiB). Its SHA256 and provenance are in [`checkpoint_info.json`](checkpoint_info.json).
44
+
45
+ ## Encode frames
46
+
47
+ Run from the repository directory:
48
+
49
+ ```python
50
+ import torch
51
+ from motion_encoder import MotionEncoder
52
+
53
+ model = MotionEncoder.from_pretrained(device="cpu") # or device="cuda"
54
+ frames = torch.rand(1, 3, 256, 256).to(
55
+ device=next(model.parameters()).device,
56
+ dtype=next(model.parameters()).dtype,
57
+ )
58
+ with torch.inference_mode():
59
+ tokens, metrics = model.encode(frames)
60
+ print(tokens.shape) # [1, 12, 1, 32]
61
+ ```
62
+
63
+ `frames` must be RGB, normalized to `[0, 1]`, with shape `[N, 3, 256, 256]`.
64
+ The training preprocessing pads portrait frames horizontally to a square and resizes them to 256×256 with bilinear interpolation (`align_corners=False`). The original `HW_encoder_2` preprocessing class is included in `encoder_blocks.py`.
65
+
66
+ Each frame produces 32 tokens of 12 dimensions. The training integration flattens them into 384 dimensions per frame before retargeting. The encoder processes frames independently; temporal retargeting is outside this release.
67
+
68
+ ## Architecture
69
+
70
+ - TiTokEncoder: 24 Transformer layers, hidden width 1024, 16 attention heads.
71
+ - Patch size 16; 32 learned latent tokens.
72
+ - 12-dimensional output projection and a 4096-entry vector-quantization codebook.
73
+ - Original `is_legacy=True` token reshape is retained for checkpoint compatibility.
74
+
75
+ Implementation: `motion_encoder.py` assembles the modules; `encoder_blocks.py` contains the original encoder and preprocessing; `quantizer.py` contains the original VQ implementation.
76
+
77
+ ## Checkpoint provenance and verification
78
+
79
+ The selected checkpoint is `train_dit_5C_v6_part5/step-12200.safetensors`, dated 2025-11-28 UTC by file modification time. It is the newest checkpoint in the inspected local runs, rather than a claim of best quality or a verified paper-final checkpoint.
80
+
81
+ Training saved only trainable parameters. The selected checkpoint supplies 300 encoder/latent-token tensors. The frozen VQ codebook is restored from `train_motion_only_full_3C_20joint/step-3700.safetensors`, following the available training initialization code. This yields a complete 301-tensor encoding module. The historical run's frozen state has not been independently verified.
82
+
83
+ Validation checked tensor byte hashes against their source checkpoints, strict state-dict loading, and exact single-frame output agreement with the available original TiTok implementation. Full video-generation quality was not evaluated in this export.
84
+
85
+ ## Acknowledgments and license
86
+
87
+ This encoder builds on [TiTok / 1d-tokenizer](https://github.com/bytedance/1d-tokenizer). Source attribution is preserved. See [LICENSE](LICENSE) and [NOTICE](NOTICE).
checkpoint_info.json ADDED
@@ -0,0 +1,24 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "checkpoint": "train_dit_5C_v6_part5/step-12200.safetensors",
3
+ "checkpoint_mtime_utc": "2025-11-28T03:50:08Z",
4
+ "selection": "Latest file modification time among the inspected local checkpoints; not a quality ranking.",
5
+ "weight_file": "motion_encoder_latest.safetensors",
6
+ "weight_bytes": 607017208,
7
+ "sha256": "008316ed723d89a5cc4b5d6b2917509b67b434fdba92c722b36116bc0aa32c10",
8
+ "saved_tensors_from_selected_checkpoint": 300,
9
+ "frozen_codebook_source": "train_motion_only_full_3C_20joint/step-3700.safetensors",
10
+ "total_exported_tensors": 301,
11
+ "limitations": "Frozen VQ codebook reconstructed according to the available training script; historical runtime state has not been independently verified.",
12
+ "verification": {
13
+ "tensor_hashes_verified": 301,
14
+ "strict_load": true,
15
+ "original_forward_exact_match": true,
16
+ "output_shape": [
17
+ 1,
18
+ 12,
19
+ 1,
20
+ 32
21
+ ],
22
+ "torch": "2.7.1+cu118"
23
+ }
24
+ }
config.yaml ADDED
@@ -0,0 +1,15 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ model:
2
+ vq_model:
3
+ codebook_size: 4096
4
+ token_size: 12
5
+ use_l2_norm: true
6
+ commitment_cost: 0.25
7
+ vit_enc_model_size: large
8
+ vit_enc_patch_size: 16
9
+ num_latent_tokens: 32
10
+ quantize_mode: vq
11
+ is_legacy: true
12
+ dataset:
13
+ preprocessing:
14
+ height_size: 256
15
+ width_size: 256
encoder_blocks.py ADDED
@@ -0,0 +1,197 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Extracted for IM-Animation: encoder and preprocessing only; original forward logic retained.
2
+ """Building blocks for TiTok.
3
+
4
+ Copyright (2024) Bytedance Ltd. and/or its affiliates
5
+
6
+ Licensed under the Apache License, Version 2.0 (the "License");
7
+ you may not use this file except in compliance with the License.
8
+ You may obtain a copy of the License at
9
+
10
+ http://www.apache.org/licenses/LICENSE-2.0
11
+
12
+ Unless required by applicable law or agreed to in writing, software
13
+ distributed under the License is distributed on an "AS IS" BASIS,
14
+ WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
15
+ See the License for the specific language governing permissions and
16
+ limitations under the License.
17
+
18
+ Reference:
19
+ https://github.com/mlfoundations/open_clip/blob/main/src/open_clip/transformer.py
20
+ https://github.com/baofff/U-ViT/blob/main/libs/timm.py
21
+ """
22
+
23
+ import torch
24
+ import torch.nn as nn
25
+ import torch.nn.functional as F
26
+ import torch.utils.checkpoint
27
+ from collections import OrderedDict
28
+ from einops import rearrange
29
+
30
+ class ResidualAttentionBlock(nn.Module):
31
+ def __init__(
32
+ self,
33
+ d_model,
34
+ n_head,
35
+ mlp_ratio = 4.0,
36
+ act_layer = nn.GELU,
37
+ norm_layer = nn.LayerNorm
38
+ ):
39
+ super().__init__()
40
+
41
+ self.ln_1 = norm_layer(d_model)
42
+ self.attn = nn.MultiheadAttention(d_model, n_head)
43
+ self.mlp_ratio = mlp_ratio
44
+ # optionally we can disable the FFN
45
+ if mlp_ratio > 0:
46
+ self.ln_2 = norm_layer(d_model)
47
+ mlp_width = int(d_model * mlp_ratio)
48
+ self.mlp = nn.Sequential(OrderedDict([
49
+ ("c_fc", nn.Linear(d_model, mlp_width)),
50
+ ("gelu", act_layer()),
51
+ ("c_proj", nn.Linear(mlp_width, d_model))
52
+ ]))
53
+
54
+ def attention(
55
+ self,
56
+ x: torch.Tensor
57
+ ):
58
+ return self.attn(x, x, x, need_weights=False)[0]
59
+
60
+ def forward(
61
+ self,
62
+ x: torch.Tensor,
63
+ ):
64
+ attn_output = self.attention(x=self.ln_1(x))
65
+ x = x + attn_output
66
+ if self.mlp_ratio > 0:
67
+ x = x + self.mlp(self.ln_2(x))
68
+ return x
69
+
70
+ def _expand_token(token, batch_size: int):
71
+ return token.unsqueeze(0).expand(batch_size, -1, -1)
72
+
73
+ class TiTokEncoder(nn.Module):
74
+ def __init__(self, config):
75
+ super().__init__()
76
+ self.config = config
77
+ self.width_size = config.dataset.preprocessing.width_size
78
+ self.height_size = config.dataset.preprocessing.height_size
79
+ self.patch_size = config.model.vq_model.vit_enc_patch_size
80
+ self.grid_size_w = self.width_size // self.patch_size
81
+ self.grid_size_h = self.height_size // self.patch_size
82
+ self.model_size = config.model.vq_model.vit_enc_model_size
83
+ self.num_latent_tokens = config.model.vq_model.num_latent_tokens
84
+ self.token_size = config.model.vq_model.token_size
85
+
86
+ if config.model.vq_model.get("quantize_mode", "vq") == "vae":
87
+ self.token_size = self.token_size * 2 # needs to split into mean and std
88
+
89
+ self.is_legacy = config.model.vq_model.get("is_legacy", True)
90
+
91
+ self.width = {
92
+ "small": 512,
93
+ "base": 768,
94
+ "large": 1024,
95
+ }[self.model_size]
96
+ self.num_layers = {
97
+ "small": 8,
98
+ "base": 12,
99
+ "large": 24,
100
+ }[self.model_size]
101
+ self.num_heads = {
102
+ "small": 8,
103
+ "base": 12,
104
+ "large": 16,
105
+ }[self.model_size]
106
+
107
+ self.patch_embed = nn.Conv2d(
108
+ in_channels=3, out_channels=self.width,
109
+ kernel_size=self.patch_size, stride=self.patch_size,padding = (4,2), bias=True)
110
+
111
+ scale = self.width ** -0.5
112
+ self.class_embedding = nn.Parameter(scale * torch.randn(1, self.width))
113
+ self.positional_embedding = nn.Parameter(
114
+ scale * torch.randn(self.grid_size_h*self.grid_size_w + 1, self.width))
115
+ self.latent_token_positional_embedding = nn.Parameter(
116
+ scale * torch.randn(self.num_latent_tokens, self.width))
117
+ self.ln_pre = nn.LayerNorm(self.width)
118
+ self.transformer = nn.ModuleList()
119
+ for i in range(self.num_layers):
120
+ self.transformer.append(ResidualAttentionBlock(
121
+ self.width, self.num_heads, mlp_ratio=4.0
122
+ ))
123
+ self.ln_post = nn.LayerNorm(self.width)
124
+ self.conv_out = nn.Conv2d(self.width, self.token_size, kernel_size=1, bias=True)
125
+
126
+ def forward(self, pixel_values, latent_tokens):
127
+ batch_size = pixel_values.shape[0]
128
+ x = pixel_values
129
+ x = self.patch_embed(x)
130
+ x = x.reshape(x.shape[0], x.shape[1], -1)
131
+ x = x.permute(0, 2, 1) # shape = [*, grid ** 2, width]
132
+ # class embeddings and positional embeddings
133
+ x = torch.cat([_expand_token(self.class_embedding, x.shape[0]).to(x.dtype), x], dim=1)
134
+ x = x + self.positional_embedding.to(x.dtype) # shape = [*, grid ** 2 + 1, width]
135
+
136
+
137
+ latent_tokens = _expand_token(latent_tokens, x.shape[0]).to(x.dtype)
138
+ latent_tokens = latent_tokens + self.latent_token_positional_embedding.to(x.dtype)
139
+ x = torch.cat([x, latent_tokens], dim=1)
140
+ def create_custom_forward(module):
141
+ def custom_forward(*inputs):
142
+ return module(*inputs)
143
+ return custom_forward
144
+ x = self.ln_pre(x)
145
+ x = x.permute(1, 0, 2) # NLD -> LND
146
+ for i in range(self.num_layers):
147
+ # x = self.transformer[i](x)
148
+ #with torch.autograd.graph.save_on_cpu():
149
+ x = torch.utils.checkpoint.checkpoint(create_custom_forward(self.transformer[i]), x,use_reentrant=False)
150
+
151
+ x = x.permute(1, 0, 2) # LND -> NLD
152
+
153
+ latent_tokens = x[:, 1+self.grid_size_h*self.grid_size_w:]
154
+ latent_tokens = self.ln_post(latent_tokens)
155
+ # fake 2D shape
156
+ if self.is_legacy:
157
+ latent_tokens = latent_tokens.reshape(batch_size, self.width, self.num_latent_tokens, 1)
158
+ else:
159
+ # Fix legacy problem.
160
+ latent_tokens = latent_tokens.reshape(batch_size, self.num_latent_tokens, self.width, 1).permute(0, 2, 1, 3)
161
+ latent_tokens = self.conv_out(latent_tokens)
162
+ latent_tokens = latent_tokens.reshape(batch_size, self.token_size, 1, self.num_latent_tokens)
163
+ return latent_tokens
164
+
165
+ class HW_encoder_2(nn.Module):
166
+ def __init__(self, in_channels):
167
+ super(HW_encoder_2, self).__init__()
168
+ # self.conv0 = nn.Conv2d(in_channels, in_channels, kernel_size=3, padding=1)
169
+
170
+ # self.conv1 = nn.Conv2d(in_channels*4, in_channels, kernel_size=3, padding=1)
171
+ # self.conv2 = nn.Conv2d(in_channels*4 , in_channels, kernel_size=3, padding=1)
172
+ # self.conv3 = nn.Conv2d(in_channels*4 , in_channels , kernel_size=3, padding=1)
173
+ # # self.conv4 = nn.Conv2d(in_channels , in_channels//4 , kernel_size=3, padding=1)
174
+
175
+ # def pixel_shuffle(self, x, scale_factor=0.5):
176
+ # n, c, h, w = x.size()
177
+ # new_h = int(h * scale_factor)
178
+ # new_w = int(w * scale_factor)
179
+ # x = x.view(n, int(c / (scale_factor ** 2)), new_h, new_w)
180
+
181
+ # return x
182
+ def forward(self, x):
183
+ B, C, T, H, W = x.shape
184
+ x = rearrange(x, "b c f h w -> (b f) c h w")
185
+
186
+ # Step 1: Pad the width from 480 to 832
187
+ padding_width = (H - W) // 2
188
+ x = F.pad(x, (padding_width, padding_width, 0, 0)) # Pad width only
189
+
190
+ # Step 2: Resize to target width 256
191
+ target_size = (256, 256)
192
+ x = F.interpolate(x, size=target_size, mode='bilinear', align_corners=False)
193
+
194
+ # Rearrange back to original shape
195
+ x = rearrange(x, "(b f) c h w -> b c f h w", f=T)
196
+
197
+ return x
motion_encoder.py ADDED
@@ -0,0 +1,38 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Encoder-only assembly of the original TiTok motion encoding path.
2
+
3
+ TiTokEncoder, VectorQuantizer, HW_encoder_2 retain the original implementations.
4
+ Input to encode: preprocessed frames [N, 3, 256, 256], RGB in [0, 1].
5
+ """
6
+ from pathlib import Path
7
+ import torch
8
+ from torch import nn
9
+ from omegaconf import OmegaConf
10
+ from safetensors.torch import load_file
11
+ from encoder_blocks import TiTokEncoder, HW_encoder_2
12
+ from quantizer import VectorQuantizer
13
+
14
+ class MotionEncoder(nn.Module):
15
+ def __init__(self, config):
16
+ super().__init__()
17
+ self.encoder = TiTokEncoder(config)
18
+ vq = config.model.vq_model
19
+ self.latent_tokens = nn.Parameter(torch.empty(vq.num_latent_tokens, self.encoder.width))
20
+ self.quantize = VectorQuantizer(codebook_size=vq.codebook_size,
21
+ token_size=vq.token_size, commitment_cost=vq.commitment_cost,
22
+ use_l2_norm=vq.use_l2_norm)
23
+
24
+ def encode(self, frames):
25
+ z = self.encoder(pixel_values=frames, latent_tokens=self.latent_tokens)
26
+ return self.quantize(z)
27
+
28
+ def forward(self, frames):
29
+ return self.encode(frames)
30
+
31
+ @classmethod
32
+ def from_pretrained(cls, directory=None, device='cpu'):
33
+ directory = Path(directory or Path(__file__).parent)
34
+ config = OmegaConf.load(directory / 'config.yaml')
35
+ with torch.device('meta'):
36
+ model = cls(config)
37
+ model.load_state_dict(load_file(str(directory / 'motion_encoder_latest.safetensors')), strict=True, assign=True)
38
+ return model.to(device).eval()
motion_encoder_latest.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:008316ed723d89a5cc4b5d6b2917509b67b434fdba92c722b36116bc0aa32c10
3
+ size 607017208
quantizer.py ADDED
@@ -0,0 +1,170 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Vector quantizer.
2
+
3
+ Copyright (2024) Bytedance Ltd. and/or its affiliates
4
+
5
+ Licensed under the Apache License, Version 2.0 (the "License");
6
+ you may not use this file except in compliance with the License.
7
+ You may obtain a copy of the License at
8
+
9
+ http://www.apache.org/licenses/LICENSE-2.0
10
+
11
+ Unless required by applicable law or agreed to in writing, software
12
+ distributed under the License is distributed on an "AS IS" BASIS,
13
+ WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
14
+ See the License for the specific language governing permissions and
15
+ limitations under the License.
16
+
17
+ Reference:
18
+ https://github.com/CompVis/taming-transformers/blob/master/taming/modules/vqvae/quantize.py
19
+ https://github.com/google-research/magvit/blob/main/videogvt/models/vqvae.py
20
+ https://github.com/CompVis/latent-diffusion/blob/main/ldm/modules/distributions/distributions.py
21
+ https://github.com/lyndonzheng/CVQ-VAE/blob/main/quantise.py
22
+ """
23
+ from typing import Mapping, Text, Tuple
24
+
25
+ import torch
26
+ from einops import rearrange
27
+ from accelerate.utils.operations import gather
28
+ from torch.cuda.amp import autocast
29
+
30
+ class VectorQuantizer(torch.nn.Module):
31
+ def __init__(self,
32
+ codebook_size: int = 1024,
33
+ token_size: int = 256,
34
+ commitment_cost: float = 0.25,
35
+ use_l2_norm: bool = False,
36
+ clustering_vq: bool = False
37
+ ):
38
+ super().__init__()
39
+ self.codebook_size = codebook_size
40
+ self.token_size = token_size
41
+ self.commitment_cost = commitment_cost
42
+
43
+ self.embedding = torch.nn.Embedding(codebook_size, token_size)
44
+ self.embedding.weight.data.uniform_(-1.0 / codebook_size, 1.0 / codebook_size)
45
+ self.use_l2_norm = use_l2_norm
46
+
47
+ self.clustering_vq = clustering_vq
48
+ if clustering_vq:
49
+ self.decay = 0.99
50
+ self.register_buffer("embed_prob", torch.zeros(self.codebook_size))
51
+
52
+ # Ensure quantization is performed using f32
53
+ # @autocast(enabled=False)
54
+ def forward(self, z: torch.Tensor) -> Tuple[torch.Tensor, Mapping[Text, torch.Tensor]]:
55
+ # z = z.float()
56
+ z = rearrange(z, 'b c h w -> b h w c').contiguous()
57
+ z_flattened = rearrange(z, 'b h w c -> (b h w) c')
58
+ unnormed_z_flattened = z_flattened
59
+
60
+ if self.use_l2_norm:
61
+ z_flattened = torch.nn.functional.normalize(z_flattened, dim=-1)
62
+ embedding = torch.nn.functional.normalize(self.embedding.weight, dim=-1)
63
+ else:
64
+ embedding = self.embedding.weight
65
+ d = torch.sum(z_flattened**2, dim=1, keepdim=True) + \
66
+ torch.sum(embedding**2, dim=1) - 2 * \
67
+ torch.einsum('bd,dn->bn', z_flattened, embedding.T)
68
+
69
+ min_encoding_indices = torch.argmin(d, dim=1) # num_ele
70
+ z_quantized = self.get_codebook_entry(min_encoding_indices).view(z.shape)
71
+
72
+ if self.use_l2_norm:
73
+ z = torch.nn.functional.normalize(z, dim=-1)
74
+
75
+ # compute loss for embedding
76
+ commitment_loss = self.commitment_cost * torch.mean((z_quantized.detach() - z) **2)
77
+ codebook_loss = torch.mean((z_quantized - z.detach()) **2)
78
+
79
+ if self.clustering_vq and self.training:
80
+ with torch.no_grad():
81
+ # Gather distance matrix from all GPUs.
82
+ encoding_indices = gather(min_encoding_indices)
83
+ if len(min_encoding_indices.shape) != 1:
84
+ raise ValueError(f"min_encoding_indices in a wrong shape, {min_encoding_indices.shape}")
85
+ # Compute and update the usage of each entry in the codebook.
86
+ encodings = torch.zeros(encoding_indices.shape[0], self.codebook_size, device=z.device)
87
+ encodings.scatter_(1, encoding_indices.unsqueeze(1), 1)
88
+ avg_probs = torch.mean(encodings, dim=0)
89
+ self.embed_prob.mul_(self.decay).add_(avg_probs, alpha=1-self.decay)
90
+ # Closest sampling to update the codebook.
91
+ all_d = gather(d)
92
+ all_unnormed_z_flattened = gather(unnormed_z_flattened).detach()
93
+ if all_d.shape[0] != all_unnormed_z_flattened.shape[0]:
94
+ raise ValueError(
95
+ "all_d and all_unnormed_z_flattened have different length" +
96
+ f"{all_d.shape}, {all_unnormed_z_flattened.shape}")
97
+ indices = torch.argmin(all_d, dim=0)
98
+ random_feat = all_unnormed_z_flattened[indices]
99
+ # Decay parameter based on the average usage.
100
+ decay = torch.exp(-(self.embed_prob * self.codebook_size * 10) /
101
+ (1 - self.decay) - 1e-3).unsqueeze(1).repeat(1, self.token_size)
102
+ self.embedding.weight.data = self.embedding.weight.data * (1 - decay) + random_feat * decay
103
+
104
+ loss = commitment_loss + codebook_loss
105
+
106
+ # preserve gradients
107
+ z_quantized = z + (z_quantized - z).detach()
108
+
109
+ # reshape back to match original input shape
110
+ z_quantized = rearrange(z_quantized, 'b h w c -> b c h w').contiguous()
111
+
112
+ result_dict = dict(
113
+ quantizer_loss=loss,
114
+ commitment_loss=commitment_loss,
115
+ codebook_loss=codebook_loss,
116
+ min_encoding_indices=min_encoding_indices.view(z_quantized.shape[0], z_quantized.shape[2], z_quantized.shape[3])
117
+ )
118
+
119
+ return z_quantized, result_dict
120
+
121
+ def get_codebook_entry(self, indices):
122
+ if len(indices.shape) == 1:
123
+ z_quantized = self.embedding(indices)
124
+ elif len(indices.shape) == 2:
125
+ z_quantized = torch.einsum('bd,dn->bn', indices, self.embedding.weight)
126
+ else:
127
+ raise NotImplementedError
128
+ if self.use_l2_norm:
129
+ z_quantized = torch.nn.functional.normalize(z_quantized, dim=-1)
130
+ return z_quantized
131
+
132
+
133
+ class DiagonalGaussianDistribution(object):
134
+ @autocast(enabled=False)
135
+ def __init__(self, parameters, deterministic=False):
136
+ """Initializes a Gaussian distribution instance given the parameters.
137
+
138
+ Args:
139
+ parameters (torch.Tensor): The parameters for the Gaussian distribution. It is expected
140
+ to be in shape [B, 2 * C, *], where B is batch size, and C is the embedding dimension.
141
+ First C channels are used for mean and last C are used for logvar in the Gaussian distribution.
142
+ deterministic (bool): Whether to use deterministic sampling. When it is true, the sampling results
143
+ is purely based on mean (i.e., std = 0).
144
+ """
145
+ self.parameters = parameters
146
+ self.mean, self.logvar = torch.chunk(parameters.float(), 2, dim=1)
147
+ self.logvar = torch.clamp(self.logvar, -30.0, 20.0)
148
+ self.deterministic = deterministic
149
+ self.std = torch.exp(0.5 * self.logvar)
150
+ self.var = torch.exp(self.logvar)
151
+ if self.deterministic:
152
+ self.var = self.std = torch.zeros_like(self.mean).to(device=self.parameters.device)
153
+
154
+ @autocast(enabled=False)
155
+ def sample(self):
156
+ x = self.mean.float() + self.std.float() * torch.randn(self.mean.shape).to(device=self.parameters.device)
157
+ return x
158
+
159
+ @autocast(enabled=False)
160
+ def mode(self):
161
+ return self.mean
162
+
163
+ @autocast(enabled=False)
164
+ def kl(self):
165
+ if self.deterministic:
166
+ return torch.Tensor([0.])
167
+ else:
168
+ return 0.5 * torch.sum(torch.pow(self.mean.float(), 2)
169
+ + self.var.float() - 1.0 - self.logvar.float(),
170
+ dim=[1, 2])
requirements.txt ADDED
@@ -0,0 +1,6 @@
 
 
 
 
 
 
 
1
+ torch>=2.1
2
+ einops>=0.7
3
+ omegaconf>=2.3
4
+ safetensors>=0.4
5
+ accelerate>=0.26
6
+ huggingface_hub>=0.24