tijayantML commited on
Commit
1d7ff26
·
verified ·
1 Parent(s): 910142f

Add flow and pixel IDM weights (Apache 2.0)

Browse files
LICENSE ADDED
@@ -0,0 +1,202 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+
2
+ Apache License
3
+ Version 2.0, January 2004
4
+ http://www.apache.org/licenses/
5
+
6
+ TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
7
+
8
+ 1. Definitions.
9
+
10
+ "License" shall mean the terms and conditions for use, reproduction,
11
+ and distribution as defined by Sections 1 through 9 of this document.
12
+
13
+ "Licensor" shall mean the copyright owner or entity authorized by
14
+ the copyright owner that is granting the License.
15
+
16
+ "Legal Entity" shall mean the union of the acting entity and all
17
+ other entities that control, are controlled by, or are under common
18
+ control with that entity. For the purposes of this definition,
19
+ "control" means (i) the power, direct or indirect, to cause the
20
+ direction or management of such entity, whether by contract or
21
+ otherwise, or (ii) ownership of fifty percent (50%) or more of the
22
+ outstanding shares, or (iii) beneficial ownership of such entity.
23
+
24
+ "You" (or "Your") shall mean an individual or Legal Entity
25
+ exercising permissions granted by this License.
26
+
27
+ "Source" form shall mean the preferred form for making modifications,
28
+ including but not limited to software source code, documentation
29
+ source, and configuration files.
30
+
31
+ "Object" form shall mean any form resulting from mechanical
32
+ transformation or translation of a Source form, including but
33
+ not limited to compiled object code, generated documentation,
34
+ and conversions to other media types.
35
+
36
+ "Work" shall mean the work of authorship, whether in Source or
37
+ Object form, made available under the License, as indicated by a
38
+ copyright notice that is included in or attached to the work
39
+ (an example is provided in the Appendix below).
40
+
41
+ "Derivative Works" shall mean any work, whether in Source or Object
42
+ form, that is based on (or derived from) the Work and for which the
43
+ editorial revisions, annotations, elaborations, or other modifications
44
+ represent, as a whole, an original work of authorship. For the purposes
45
+ of this License, Derivative Works shall not include works that remain
46
+ separable from, or merely link (or bind by name) to the interfaces of,
47
+ the Work and Derivative Works thereof.
48
+
49
+ "Contribution" shall mean any work of authorship, including
50
+ the original version of the Work and any modifications or additions
51
+ to that Work or Derivative Works thereof, that is intentionally
52
+ submitted to Licensor for inclusion in the Work by the copyright owner
53
+ or by an individual or Legal Entity authorized to submit on behalf of
54
+ the copyright owner. For the purposes of this definition, "submitted"
55
+ means any form of electronic, verbal, or written communication sent
56
+ to the Licensor or its representatives, including but not limited to
57
+ communication on electronic mailing lists, source code control systems,
58
+ and issue tracking systems that are managed by, or on behalf of, the
59
+ Licensor for the purpose of discussing and improving the Work, but
60
+ excluding communication that is conspicuously marked or otherwise
61
+ designated in writing by the copyright owner as "Not a Contribution."
62
+
63
+ "Contributor" shall mean Licensor and any individual or Legal Entity
64
+ on behalf of whom a Contribution has been received by Licensor and
65
+ subsequently incorporated within the Work.
66
+
67
+ 2. Grant of Copyright License. Subject to the terms and conditions of
68
+ this License, each Contributor hereby grants to You a perpetual,
69
+ worldwide, non-exclusive, no-charge, royalty-free, irrevocable
70
+ copyright license to reproduce, prepare Derivative Works of,
71
+ publicly display, publicly perform, sublicense, and distribute the
72
+ Work and such Derivative Works in Source or Object form.
73
+
74
+ 3. Grant of Patent License. Subject to the terms and conditions of
75
+ this License, each Contributor hereby grants to You a perpetual,
76
+ worldwide, non-exclusive, no-charge, royalty-free, irrevocable
77
+ (except as stated in this section) patent license to make, have made,
78
+ use, offer to sell, sell, import, and otherwise transfer the Work,
79
+ where such license applies only to those patent claims licensable
80
+ by such Contributor that are necessarily infringed by their
81
+ Contribution(s) alone or by combination of their Contribution(s)
82
+ with the Work to which such Contribution(s) was submitted. If You
83
+ institute patent litigation against any entity (including a
84
+ cross-claim or counterclaim in a lawsuit) alleging that the Work
85
+ or a Contribution incorporated within the Work constitutes direct
86
+ or contributory patent infringement, then any patent licenses
87
+ granted to You under this License for that Work shall terminate
88
+ as of the date such litigation is filed.
89
+
90
+ 4. Redistribution. You may reproduce and distribute copies of the
91
+ Work or Derivative Works thereof in any medium, with or without
92
+ modifications, and in Source or Object form, provided that You
93
+ meet the following conditions:
94
+
95
+ (a) You must give any other recipients of the Work or
96
+ Derivative Works a copy of this License; and
97
+
98
+ (b) You must cause any modified files to carry prominent notices
99
+ stating that You changed the files; and
100
+
101
+ (c) You must retain, in the Source form of any Derivative Works
102
+ that You distribute, all copyright, patent, trademark, and
103
+ attribution notices from the Source form of the Work,
104
+ excluding those notices that do not pertain to any part of
105
+ the Derivative Works; and
106
+
107
+ (d) If the Work includes a "NOTICE" text file as part of its
108
+ distribution, then any Derivative Works that You distribute must
109
+ include a readable copy of the attribution notices contained
110
+ within such NOTICE file, excluding those notices that do not
111
+ pertain to any part of the Derivative Works, in at least one
112
+ of the following places: within a NOTICE text file distributed
113
+ as part of the Derivative Works; within the Source form or
114
+ documentation, if provided along with the Derivative Works; or,
115
+ within a display generated by the Derivative Works, if and
116
+ wherever such third-party notices normally appear. The contents
117
+ of the NOTICE file are for informational purposes only and
118
+ do not modify the License. You may add Your own attribution
119
+ notices within Derivative Works that You distribute, alongside
120
+ or as an addendum to the NOTICE text from the Work, provided
121
+ that such additional attribution notices cannot be construed
122
+ as modifying the License.
123
+
124
+ You may add Your own copyright statement to Your modifications and
125
+ may provide additional or different license terms and conditions
126
+ for use, reproduction, or distribution of Your modifications, or
127
+ for any such Derivative Works as a whole, provided Your use,
128
+ reproduction, and distribution of the Work otherwise complies with
129
+ the conditions stated in this License.
130
+
131
+ 5. Submission of Contributions. Unless You explicitly state otherwise,
132
+ any Contribution intentionally submitted for inclusion in the Work
133
+ by You to the Licensor shall be under the terms and conditions of
134
+ this License, without any additional terms or conditions.
135
+ Notwithstanding the above, nothing herein shall supersede or modify
136
+ the terms of any separate license agreement you may have executed
137
+ with Licensor regarding such Contributions.
138
+
139
+ 6. Trademarks. This License does not grant permission to use the trade
140
+ names, trademarks, service marks, or product names of the Licensor,
141
+ except as required for reasonable and customary use in describing the
142
+ origin of the Work and reproducing the content of the NOTICE file.
143
+
144
+ 7. Disclaimer of Warranty. Unless required by applicable law or
145
+ agreed to in writing, Licensor provides the Work (and each
146
+ Contributor provides its Contributions) on an "AS IS" BASIS,
147
+ WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
148
+ implied, including, without limitation, any warranties or conditions
149
+ of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
150
+ PARTICULAR PURPOSE. You are solely responsible for determining the
151
+ appropriateness of using or redistributing the Work and assume any
152
+ risks associated with Your exercise of permissions under this License.
153
+
154
+ 8. Limitation of Liability. In no event and under no legal theory,
155
+ whether in tort (including negligence), contract, or otherwise,
156
+ unless required by applicable law (such as deliberate and grossly
157
+ negligent acts) or agreed to in writing, shall any Contributor be
158
+ liable to You for damages, including any direct, indirect, special,
159
+ incidental, or consequential damages of any character arising as a
160
+ result of this License or out of the use or inability to use the
161
+ Work (including but not limited to damages for loss of goodwill,
162
+ work stoppage, computer failure or malfunction, or any and all
163
+ other commercial damages or losses), even if such Contributor
164
+ has been advised of the possibility of such damages.
165
+
166
+ 9. Accepting Warranty or Additional Liability. While redistributing
167
+ the Work or Derivative Works thereof, You may choose to offer,
168
+ and charge a fee for, acceptance of support, warranty, indemnity,
169
+ or other liability obligations and/or rights consistent with this
170
+ License. However, in accepting such obligations, You may act only
171
+ on Your own behalf and on Your sole responsibility, not on behalf
172
+ of any other Contributor, and only if You agree to indemnify,
173
+ defend, and hold each Contributor harmless for any liability
174
+ incurred by, or claims asserted against, such Contributor by reason
175
+ of your accepting any such warranty or additional liability.
176
+
177
+ END OF TERMS AND CONDITIONS
178
+
179
+ APPENDIX: How to apply the Apache License to your work.
180
+
181
+ To apply the Apache License to your work, attach the following
182
+ boilerplate notice, with the fields enclosed by brackets "[]"
183
+ replaced with your own identifying information. (Don't include
184
+ the brackets!) The text should be enclosed in the appropriate
185
+ comment syntax for the file format. We also recommend that a
186
+ file or class name and description of purpose be included on the
187
+ same "printed page" as the copyright notice for easier
188
+ identification within third-party archives.
189
+
190
+ Copyright [yyyy] [name of copyright owner]
191
+
192
+ Licensed under the Apache License, Version 2.0 (the "License");
193
+ you may not use this file except in compliance with the License.
194
+ You may obtain a copy of the License at
195
+
196
+ http://www.apache.org/licenses/LICENSE-2.0
197
+
198
+ Unless required by applicable law or agreed to in writing, software
199
+ distributed under the License is distributed on an "AS IS" BASIS,
200
+ WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
201
+ See the License for the specific language governing permissions and
202
+ limitations under the License.
NOTICE ADDED
@@ -0,0 +1,8 @@
 
 
 
 
 
 
 
 
 
1
+ IDM inverse dynamics models
2
+ Copyright 2026 Reka AI
3
+
4
+ The weights and code in this repository use the Apache License 2.0 (see LICENSE).
5
+
6
+ RAFT-small optical flow weights: BSD-3-Clause, from torchvision
7
+ (torchvision.models.optical_flow.raft_small). flow/inference.py downloads them at
8
+ run time. They are not part of this repository and are not redistributed here.
README.md ADDED
@@ -0,0 +1,42 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ license: apache-2.0
3
+ library_name: pytorch
4
+ pipeline_tag: video-classification
5
+ tags:
6
+ - inverse-dynamics
7
+ - camera-motion
8
+ - optical-flow
9
+ - counter-strike-2
10
+ ---
11
+
12
+ # IDM: inverse dynamics models for camera motion
13
+
14
+ Two models predict camera motion (W/A/S/D/Shift, yaw, pitch) from video.
15
+ Both trained on Counter-Strike 2 renders.
16
+
17
+ | | [flow](flow/) | [pixel](pixel/) |
18
+ |---|---|---|
19
+ | Input | Optical flow (RAFT-small) | Raw frames |
20
+ | Parameters | 1,795,337 trained, plus 990,162 in frozen RAFT-small (2,785,499 in total) | 9,836,063 |
21
+ | Output | One prediction per 17-frame window | One prediction per frame |
22
+ | Real-video turn accuracy | 91.5 | 51.1 † |
23
+ | Counter-Strike 2 turn accuracy | 84.9 | 70.9 † |
24
+ | Dependencies | torch, torchvision, av, opencv, safetensors | torch, av, opencv, safetensors |
25
+
26
+ † The number comes from a write-up and has no run log yet.
27
+
28
+ ```text
29
+ idm-hf/
30
+ flow/ model.safetensors config.json inference.py README.md
31
+ pixel/ model.safetensors config.json inference.py README.md
32
+ ```
33
+
34
+ Each folder runs alone: `python inference.py clip.mp4`.
35
+
36
+ Status: preliminary release. The scores marked † are unconfirmed.
37
+
38
+ ## License
39
+
40
+ - The weights and the code in this repository use the Apache License 2.0. See `LICENSE`.
41
+ - RAFT-small weights use the BSD-3 license. They come from torchvision and are not in this repository. See `NOTICE`.
42
+ - The models trained on Counter-Strike 2 gameplay captures. This repository holds no game assets, no clips and no training data.
flow/README.md ADDED
@@ -0,0 +1,77 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ license: apache-2.0
3
+ library_name: pytorch
4
+ pipeline_tag: video-classification
5
+ tags: [inverse-dynamics, camera-motion, optical-flow, counter-strike-2]
6
+ ---
7
+
8
+ # IDM flow model (flow_transformer_phase5_v4)
9
+
10
+ This model predicts camera motion from video. It reads optical flow, not pixels.
11
+
12
+ - Keys: W, A, S, D, Shift (one probability each).
13
+ - Rotation: yaw and pitch in degrees.
14
+
15
+ Status: preliminary release. The scores marked † are unconfirmed.
16
+
17
+ ## How it works
18
+
19
+ ```text
20
+ video -> 5 frames (every 4th) -> RAFT-small flow x4 pairs -> FlowTransformer -> keys + yaw/pitch
21
+ 128x128 frozen, downloaded 1.80 M params
22
+ ```
23
+
24
+ - RAFT-small is a frozen torchvision model. Its weights are not in this folder.
25
+ The script downloads them on first run (BSD-3 license).
26
+ - The FlowTransformer has 1,795,337 parameters. It has 4 layers, width 192 and 4 heads.
27
+ - One window is 17 source frames (5 sampled frames, gap 4). Each window gives one prediction.
28
+
29
+ ## Use
30
+
31
+ ```bash
32
+ pip install torch torchvision av opencv-python safetensors numpy
33
+ python inference.py clip.mp4 --stride 17
34
+ ```
35
+
36
+ The script prints JSON. Each window has `key_probabilities`, `keys_pressed`, `yaw_deg` and `pitch_deg`.
37
+
38
+ `keys_pressed` uses one threshold per key, tuned at 5 frames: W 0.54, A 0.40, S 0.54, D 0.70, Shift 0.20.
39
+ Yaw and pitch are signed. The sign is a classifier. The size comes from a log-magnitude regressor.
40
+
41
+ ## Files
42
+
43
+ | File | Content |
44
+ |---|---|
45
+ | `model.safetensors` | FlowTransformer weights only. No optimizer state. Float32. |
46
+ | `config.json` | Architecture, thresholds, frame count, frame gap. |
47
+ | `inference.py` | Standalone script. No internal imports. |
48
+
49
+ ## Results
50
+
51
+ | Test set | Turn | Three-way |
52
+ |---|---|---|
53
+ | Real videos | 91.5 | 85.9 |
54
+ | Action-camera footage | 72.0 † | 75.0 † |
55
+ | Counter-Strike 2 | 84.9 | 84.5 |
56
+
57
+ - † means the number comes from a write-up. It has no run log yet. Treat it as unconfirmed.
58
+ - Key exact match on the held-out CS2 split (4,590 clips, 35,091 rows):
59
+ 0.7960 at threshold 0.5 and 0.8390 with the tuned thresholds.
60
+ - Best validation loss in the checkpoint: 0.8741 (epoch 10).
61
+
62
+ ## Limits
63
+
64
+ - The model trained on Counter-Strike 2 renders only.
65
+ - The default frame gap is 4 source frames, tuned on game footage. Other gaps change the result.
66
+ - Real walking is slower than game movement. On walking-tour video a gap of 12 gave a strong forward signal.
67
+ The gap comes from `frame_gap` in `config.json`. Copy the file, set the value, and pass it with `--config`.
68
+ Use `--stride` to match the new window length: (num_frames - 1) * frame_gap + 1 source frames.
69
+ - The gap embedding accepts values from 0 to 63.
70
+ - Yaw and pitch are per window, not per frame.
71
+
72
+ ## License
73
+
74
+ - The weights and `inference.py` use the Apache License 2.0. See `LICENSE` in the repository root.
75
+ - RAFT-small weights: BSD-3, from torchvision.
76
+ - The models trained on Counter-Strike 2 gameplay captures. This repository holds no game assets, no clips and no training data.
77
+ - The terms of the evaluation datasets are not checked. Do not redistribute their clips.
flow/config.json ADDED
@@ -0,0 +1,23 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "model_type": "idm-flow-transformer",
3
+ "name": "flow_transformer_phase5_v4",
4
+ "num_keys": 5,
5
+ "key_order": ["W", "A", "S", "D", "Shift"],
6
+ "key_thresholds": {"W": 0.54, "A": 0.40, "S": 0.54, "D": 0.70, "Shift": 0.20},
7
+ "thresholds_tuned_at_num_frames": 5,
8
+ "token_dim": 192,
9
+ "num_layers": 4,
10
+ "num_heads": 4,
11
+ "ffn_dim": 512,
12
+ "head_hidden": 256,
13
+ "dropout": 0.1,
14
+ "num_flow_frames": 9,
15
+ "patch_grid": 8,
16
+ "rotation_head_type": "sign_magnitude",
17
+ "resolution": 128,
18
+ "num_frames": 5,
19
+ "frame_gap": 4,
20
+ "flow_backbone": "torchvision.models.optical_flow.raft_small (Raft_Small_Weights.DEFAULT, frozen, not in this repo)",
21
+ "num_parameters": 1795337,
22
+ "training_checkpoint": {"epoch": 10, "best_val_loss": 0.8741227431995112, "run_ledger_status": "CANDIDATE"}
23
+ }
flow/inference.py ADDED
@@ -0,0 +1,190 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Flow IDM: predict W/A/S/D/Shift and yaw/pitch from a video.
2
+
3
+ Usage: python inference.py VIDEO [--weights model.safetensors] [--config config.json]
4
+
5
+ Pipeline: decode -> 128x128 RGB -> RAFT-small optical flow on consecutive sampled
6
+ frames -> FlowTransformer. Each window of (num_frames - 1) * frame_gap + 1 source
7
+ frames gives one prediction. Output is JSON on stdout.
8
+ Needs: torch, torchvision, av, opencv-python, safetensors, numpy.
9
+ """
10
+
11
+ from __future__ import annotations
12
+
13
+ import argparse
14
+ import json
15
+ from pathlib import Path
16
+
17
+ import av
18
+ import cv2
19
+ import numpy as np
20
+ import torch
21
+ import torch.nn as nn
22
+ from safetensors.torch import load_file
23
+ from torchvision.models.optical_flow import Raft_Small_Weights, raft_small
24
+
25
+ MAX_GAP_EMBED = 64
26
+
27
+
28
+ class SignMagnitudeRotationHead(nn.Module):
29
+ """Per axis: a sign logit and a log(1 + |deg|) magnitude. Output (B, 4)."""
30
+
31
+ def __init__(self, feat_dim: int, head_hidden: int, dropout: float) -> None:
32
+ super().__init__()
33
+ self.yaw_sign = self._branch(feat_dim, head_hidden, dropout)
34
+ self.yaw_mag = self._branch(feat_dim, head_hidden, dropout)
35
+ self.pitch_sign = self._branch(feat_dim, head_hidden, dropout)
36
+ self.pitch_mag = self._branch(feat_dim, head_hidden, dropout)
37
+
38
+ @staticmethod
39
+ def _branch(feat_dim: int, head_hidden: int, dropout: float) -> nn.Sequential:
40
+ return nn.Sequential(
41
+ nn.ReLU(),
42
+ nn.Linear(feat_dim, head_hidden),
43
+ nn.ReLU(inplace=True),
44
+ nn.Dropout(dropout),
45
+ nn.Linear(head_hidden, 1),
46
+ )
47
+
48
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
49
+ return torch.cat([self.yaw_sign(x), self.yaw_mag(x), self.pitch_sign(x), self.pitch_mag(x)], dim=1)
50
+
51
+
52
+ class FlowTransformer(nn.Module):
53
+ """Conv stem per flow pair -> 64 tokens per pair -> joint space-time encoder -> CLS heads."""
54
+
55
+ def __init__(self, cfg: dict) -> None:
56
+ super().__init__()
57
+ d, grid = cfg["token_dim"], cfg["patch_grid"]
58
+ self.grid = grid
59
+ self.max_pairs = cfg["num_flow_frames"]
60
+ self.stem = nn.Sequential(
61
+ nn.Conv2d(2, 64, kernel_size=4, stride=4),
62
+ nn.GroupNorm(8, 64),
63
+ nn.GELU(),
64
+ nn.Conv2d(64, 128, kernel_size=2, stride=2),
65
+ nn.GroupNorm(8, 128),
66
+ nn.GELU(),
67
+ nn.Conv2d(128, d, kernel_size=2, stride=2),
68
+ )
69
+ self.cls_token = nn.Parameter(torch.zeros(1, 1, d))
70
+ self.pos_embed = nn.Parameter(torch.zeros(1, grid * grid, d))
71
+ self.frame_embed = nn.Parameter(torch.zeros(1, self.max_pairs, d))
72
+ self.gap_embed = nn.Embedding(MAX_GAP_EMBED, d)
73
+ layer = nn.TransformerEncoderLayer(
74
+ d_model=d,
75
+ nhead=cfg["num_heads"],
76
+ dim_feedforward=cfg["ffn_dim"],
77
+ dropout=cfg["dropout"],
78
+ activation="gelu",
79
+ batch_first=True,
80
+ norm_first=True,
81
+ )
82
+ self.encoder = nn.TransformerEncoder(layer, num_layers=cfg["num_layers"], enable_nested_tensor=False)
83
+ self.norm = nn.LayerNorm(d)
84
+ self.key_head = nn.Sequential(
85
+ nn.Linear(d, cfg["head_hidden"]),
86
+ nn.GELU(),
87
+ nn.Dropout(cfg["dropout"]),
88
+ nn.Linear(cfg["head_hidden"], cfg["num_keys"]),
89
+ )
90
+ assert cfg["rotation_head_type"] == "sign_magnitude"
91
+ self.rotation_head = SignMagnitudeRotationHead(d, cfg["head_hidden"], cfg["dropout"])
92
+
93
+ def forward(self, flow: torch.Tensor, frame_gap: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
94
+ """flow: (B, 2*N, 128, 128), channels (fx, fy) interleaved per pair. frame_gap: (B,) long."""
95
+ b, c, h, w = flow.shape
96
+ n = c // 2
97
+ assert c % 2 == 0 and n <= self.max_pairs, f"bad channel count {c}"
98
+ assert h == w == self.grid * 16, f"expected {self.grid * 16}x{self.grid * 16}, got {h}x{w}"
99
+ feat = self.stem(flow.reshape(b * n, 2, h, w))
100
+ tok = feat.flatten(2).transpose(1, 2) + self.pos_embed
101
+ tok = tok.view(b, n, self.grid * self.grid, -1) + self.frame_embed[:, :n].unsqueeze(2)
102
+ cls = self.cls_token.repeat(b, 1, 1) + self.gap_embed(frame_gap.clamp(0, MAX_GAP_EMBED - 1)).unsqueeze(1)
103
+ x = self.encoder(torch.cat([cls, tok.flatten(1, 2)], dim=1))
104
+ x = self.norm(x[:, 0])
105
+ return self.key_head(x), self.rotation_head(x)
106
+
107
+
108
+ def decode_sign_magnitude(pred: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
109
+ """(B, 4) -> (yaw_deg, pitch_deg). Sign is sigmoid >= 0.5; magnitude is expm1 of the log-magnitude."""
110
+ yaw_mag = torch.expm1(pred[:, 1]).clamp_min(0.0)
111
+ pitch_mag = torch.expm1(pred[:, 3]).clamp_min(0.0)
112
+ yaw_sign = torch.where(torch.sigmoid(pred[:, 0]) >= 0.5, 1.0, -1.0)
113
+ pitch_sign = torch.where(torch.sigmoid(pred[:, 2]) >= 0.5, 1.0, -1.0)
114
+ return yaw_sign * yaw_mag, pitch_sign * pitch_mag
115
+
116
+
117
+ def load_model(weights: Path, cfg: dict) -> FlowTransformer:
118
+ model = FlowTransformer(cfg)
119
+ model.load_state_dict(load_file(str(weights)), strict=True)
120
+ return model.eval()
121
+
122
+
123
+ def decode_frames(video: Path, resolution: int) -> tuple[np.ndarray, float]:
124
+ """All frames as (N, res, res, 3) uint8 RGB, squashed to a square with INTER_AREA."""
125
+ frames: list[np.ndarray] = []
126
+ with av.open(str(video)) as container:
127
+ rate = container.streams.video[0].average_rate
128
+ fps = float(rate) if rate else 24.0
129
+ for frame in container.decode(video=0):
130
+ img = frame.to_ndarray(format="rgb24")
131
+ frames.append(cv2.resize(img, (resolution, resolution), interpolation=cv2.INTER_AREA))
132
+ assert frames, f"no frames decoded from {video}"
133
+ return np.stack(frames), fps
134
+
135
+
136
+ @torch.no_grad()
137
+ def window_flow(raft: nn.Module, transform, frames: np.ndarray, device: torch.device) -> torch.Tensor:
138
+ """(F, H, W, 3) uint8 -> (2*(F-1), H, W): RAFT flow of each consecutive pair, (fx, fy) interleaved."""
139
+ t = torch.from_numpy(frames).to(device).permute(0, 3, 1, 2).float() / 255.0
140
+ a, b = transform(t[:-1], t[1:])
141
+ flow = raft(a, b)[-1] # (F-1, 2, H, W)
142
+ return flow.reshape(-1, flow.shape[-2], flow.shape[-1])
143
+
144
+
145
+ def predict(video: Path, weights: Path, cfg: dict, device: torch.device, stride: int | None) -> dict:
146
+ model = load_model(weights, cfg).to(device)
147
+ raft = raft_small(weights=Raft_Small_Weights.DEFAULT).to(device).eval()
148
+ transform = Raft_Small_Weights.DEFAULT.transforms()
149
+ frames, fps = decode_frames(video, cfg["resolution"])
150
+ gap, n_frames = cfg["frame_gap"], cfg["num_frames"]
151
+ span = (n_frames - 1) * gap + 1
152
+ stride = stride or span
153
+ keys, thr = cfg["key_order"], cfg["key_thresholds"]
154
+ out = []
155
+ for start in range(0, len(frames) - span + 1, stride):
156
+ clip = frames[start : start + span : gap]
157
+ flow = window_flow(raft, transform, clip, device).unsqueeze(0)
158
+ key_logits, rot = model(flow, torch.tensor([gap], device=device))
159
+ prob = torch.sigmoid(key_logits)[0].cpu().tolist()
160
+ yaw, pitch = decode_sign_magnitude(rot)
161
+ out.append(
162
+ {
163
+ "start_frame": start,
164
+ "end_frame": start + span - 1,
165
+ "start_s": round(start / fps, 3),
166
+ "key_probabilities": dict(zip(keys, (round(p, 4) for p in prob))),
167
+ "keys_pressed": [k for k, p in zip(keys, prob) if p > thr[k]],
168
+ "yaw_deg": round(yaw[0].item(), 4),
169
+ "pitch_deg": round(pitch[0].item(), 4),
170
+ }
171
+ )
172
+ return {"video": str(video), "fps": fps, "num_frames": len(frames), "window_frames": span, "windows": out}
173
+
174
+
175
+ def main() -> None:
176
+ here = Path(__file__).parent
177
+ ap = argparse.ArgumentParser(description=__doc__)
178
+ ap.add_argument("video", type=Path)
179
+ ap.add_argument("--weights", type=Path, default=here / "model.safetensors")
180
+ ap.add_argument("--config", type=Path, default=here / "config.json")
181
+ ap.add_argument("--stride", type=int, default=None, help="source frames between windows (default: no overlap)")
182
+ ap.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu")
183
+ args = ap.parse_args()
184
+ cfg = json.loads(args.config.read_text())
185
+ result = predict(args.video, args.weights, cfg, torch.device(args.device), args.stride)
186
+ print(json.dumps(result, indent=2))
187
+
188
+
189
+ if __name__ == "__main__":
190
+ main()
flow/model.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:24a5dbfcf83306794ae5f1d4afb032a41f06faa5620224fe5015e36f30353ce9
3
+ size 7189300
pixel/README.md ADDED
@@ -0,0 +1,70 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ license: apache-2.0
3
+ library_name: pytorch
4
+ pipeline_tag: video-classification
5
+ tags: [inverse-dynamics, camera-motion, counter-strike-2]
6
+ ---
7
+
8
+ # IDM pixel model (Z1-106k-full)
9
+
10
+ This model predicts camera motion from raw video frames. It has no optical-flow stage.
11
+
12
+ - Keys: W, A, S, D, Shift (one probability per frame).
13
+ - Rotation: yaw and pitch in degrees, per frame.
14
+
15
+ ## How it works
16
+
17
+ ```text
18
+ video -> 128x128 RGB -> IMPALA CNN per frame -> temporal attention -> keys + yaw/pitch per frame
19
+ ImageNet norm 3 blocks, 64/128/128 2 layers, width 512
20
+ ```
21
+
22
+ - The model has 9,836,063 parameters.
23
+ - The rotation head is a hybrid: a bin classifier plus a residual.
24
+ Yaw has 13 bins. Pitch has 11 bins. The result is the bin centre plus the residual.
25
+ - The script cuts the video into 128-frame windows and averages any overlap.
26
+ A video with fewer than 16 frames gives no prediction.
27
+
28
+ ## Use
29
+
30
+ ```bash
31
+ pip install torch av opencv-python safetensors numpy
32
+ python inference.py clip.mp4
33
+ ```
34
+
35
+ The script prints JSON with one row per frame:
36
+ `key_probabilities`, `keys_pressed` (threshold 0.5), `yaw_deg` and `pitch_deg`.
37
+
38
+ ## Files
39
+
40
+ | File | Content |
41
+ |---|---|
42
+ | `model.safetensors` | Weights only. Float32. The CNN keys are stored as `cnn.blocks.*`. |
43
+ | `config.json` | Architecture and windowing. |
44
+ | `inference.py` | Standalone script. No internal imports. |
45
+
46
+ ## Results
47
+
48
+ | Test set | Turn | Three-way |
49
+ |---|---|---|
50
+ | Real videos | 51.1 † | 30.8 † |
51
+ | Action-camera footage | 76.0 † | 58.3 † |
52
+ | Counter-Strike 2 | 70.9 † | 77.5 † |
53
+
54
+ - † means the number comes from a write-up. It has no run log yet. Treat it as unconfirmed.
55
+ - The internal eval report shows about 0% on out-of-distribution real-world and marketplace footage.
56
+
57
+ ## Limits
58
+
59
+ - The model trained on Counter-Strike 2 renders only.
60
+ - It does not work on real-world footage. Use the flow model for that.
61
+ - The 128x128 resize squashes 16:9 frames to a square. The training data did the same.
62
+ - The training window was 150 frames. The reference inference code used 128. This folder uses 128.
63
+ - `num_heads` (8) is not stored in the checkpoint. It is the default of the training code.
64
+ The weights load with `strict=True`, and the output matches the original code exactly.
65
+
66
+ ## License
67
+
68
+ - The weights and `inference.py` use the Apache License 2.0. See `LICENSE` in the repository root.
69
+ - The models trained on Counter-Strike 2 gameplay captures. This repository holds no game assets, no clips and no training data.
70
+ - The terms of the evaluation datasets are not checked.
pixel/config.json ADDED
@@ -0,0 +1,23 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "model_type": "idm-impala-vpt",
3
+ "name": "Z1-106k-full",
4
+ "num_keys": 5,
5
+ "key_order": ["W", "A", "S", "D", "Shift"],
6
+ "key_threshold": 0.5,
7
+ "channels": [64, 128, 128],
8
+ "gn_groups": 8,
9
+ "spatial_pool_size": 4,
10
+ "temporal_hidden": 512,
11
+ "temporal_layers": 2,
12
+ "num_heads": 8,
13
+ "head_hidden": 256,
14
+ "dropout": 0.2,
15
+ "rotation_head_type": "hybrid",
16
+ "resolution": 128,
17
+ "window_frames": 128,
18
+ "window_stride": 128,
19
+ "min_frames": 16,
20
+ "training_window_frames": 150,
21
+ "num_parameters": 9836063,
22
+ "notes": "num_heads is not stored in the checkpoint; 8 is the default in the training code. The training window was 150 frames; the reference inference library used 128."
23
+ }
pixel/inference.py ADDED
@@ -0,0 +1,229 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Pixel IDM: predict W/A/S/D/Shift and yaw/pitch per frame from a video.
2
+
3
+ Usage: python inference.py VIDEO [--weights model.safetensors] [--config config.json]
4
+
5
+ Pipeline: decode -> 128x128 RGB (squashed) -> ImageNet normalisation -> windows of
6
+ window_frames -> per-frame IMPALA CNN + temporal attention -> per-frame heads.
7
+ Overlapping windows are averaged. Output is JSON on stdout.
8
+ Needs: torch, av, opencv-python, safetensors, numpy.
9
+ """
10
+
11
+ from __future__ import annotations
12
+
13
+ import argparse
14
+ import json
15
+ from pathlib import Path
16
+
17
+ import av
18
+ import cv2
19
+ import numpy as np
20
+ import torch
21
+ import torch.nn as nn
22
+ import torch.nn.functional as F
23
+ from safetensors.torch import load_file
24
+
25
+ YAW_BIN_CENTERS = [-70, -2.2, -0.7, -0.28, -0.12, -0.04, 0, 0.04, 0.12, 0.28, 0.7, 2.2, 70]
26
+ PITCH_BIN_CENTERS = [-16, -0.35, -0.13, -0.06, -0.02, 0, 0.02, 0.06, 0.13, 0.35, 16]
27
+ IMAGENET_MEAN = np.array([0.485, 0.456, 0.406], dtype=np.float32)
28
+ IMAGENET_STD = np.array([0.229, 0.224, 0.225], dtype=np.float32)
29
+
30
+
31
+ class ResBlock(nn.Module):
32
+ def __init__(self, ch: int, gn_groups: int) -> None:
33
+ super().__init__()
34
+ self.gn1 = nn.GroupNorm(gn_groups, ch)
35
+ self.conv1 = nn.Conv2d(ch, ch, 3, padding=1, bias=False)
36
+ self.gn2 = nn.GroupNorm(gn_groups, ch)
37
+ self.conv2 = nn.Conv2d(ch, ch, 3, padding=1, bias=False)
38
+
39
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
40
+ out = self.conv1(F.relu(self.gn1(x)))
41
+ out = self.conv2(F.relu(self.gn2(out)))
42
+ return out + x
43
+
44
+
45
+ class ImpalaBlock(nn.Module):
46
+ def __init__(self, in_ch: int, out_ch: int, gn_groups: int) -> None:
47
+ super().__init__()
48
+ self.conv = nn.Conv2d(in_ch, out_ch, 3, padding=1)
49
+ self.pool = nn.MaxPool2d(3, stride=2, padding=1)
50
+ self.res1 = ResBlock(out_ch, gn_groups)
51
+ self.res2 = ResBlock(out_ch, gn_groups)
52
+
53
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
54
+ return self.res2(self.res1(self.pool(self.conv(x))))
55
+
56
+
57
+ class ImpalaCNN(nn.Module):
58
+ def __init__(self, channels: list[int], gn_groups: int, pool: int) -> None:
59
+ super().__init__()
60
+ chs = [3, *channels]
61
+ self.blocks = nn.Sequential(*[ImpalaBlock(chs[i], chs[i + 1], gn_groups) for i in range(len(channels))])
62
+ self.pool = nn.AdaptiveAvgPool2d(pool)
63
+ self.feat_dim = channels[-1] * pool * pool
64
+
65
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
66
+ return self.pool(self.blocks(x)).flatten(1)
67
+
68
+
69
+ class PointwiseMLP(nn.Module):
70
+ def __init__(self, dim: int) -> None:
71
+ super().__init__()
72
+ self.norm = nn.LayerNorm(dim)
73
+ self.fc1 = nn.Linear(dim, dim * 4)
74
+ self.fc2 = nn.Linear(dim * 4, dim)
75
+
76
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
77
+ return x + self.fc2(F.gelu(self.fc1(self.norm(x))))
78
+
79
+
80
+ class ResidualAttentionBlock(nn.Module):
81
+ def __init__(self, dim: int, num_heads: int) -> None:
82
+ super().__init__()
83
+ self.norm = nn.LayerNorm(dim)
84
+ self.attn = nn.MultiheadAttention(dim, num_heads, batch_first=True)
85
+ self.mlp = PointwiseMLP(dim)
86
+
87
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
88
+ n = self.norm(x)
89
+ x = x + self.attn(n, n, n, need_weights=False)[0]
90
+ return self.mlp(x)
91
+
92
+
93
+ def sub_head(feat_dim: int, hidden: int, dropout: float, out_dim: int) -> nn.Sequential:
94
+ return nn.Sequential(
95
+ nn.ReLU(), nn.Linear(feat_dim, hidden), nn.ReLU(inplace=True), nn.Dropout(dropout), nn.Linear(hidden, out_dim)
96
+ )
97
+
98
+
99
+ class HybridRotationHead(nn.Module):
100
+ """Output (B, 26): yaw bin logits (13), yaw residual, pitch bin logits (11), pitch residual."""
101
+
102
+ def __init__(self, feat_dim: int, hidden: int, dropout: float) -> None:
103
+ super().__init__()
104
+ self.yaw_cls = sub_head(feat_dim, hidden, dropout, len(YAW_BIN_CENTERS))
105
+ self.yaw_res = sub_head(feat_dim, hidden, dropout, 1)
106
+ self.pitch_cls = sub_head(feat_dim, hidden, dropout, len(PITCH_BIN_CENTERS))
107
+ self.pitch_res = sub_head(feat_dim, hidden, dropout, 1)
108
+
109
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
110
+ return torch.cat([self.yaw_cls(x), self.yaw_res(x), self.pitch_cls(x), self.pitch_res(x)], dim=1)
111
+
112
+
113
+ class IDMImpalaVPT(nn.Module):
114
+ """Per-frame IMPALA CNN -> temporal attention over the window -> per-frame heads."""
115
+
116
+ def __init__(self, cfg: dict) -> None:
117
+ super().__init__()
118
+ hid = cfg["temporal_hidden"]
119
+ self.cnn = ImpalaCNN(cfg["channels"], cfg["gn_groups"], cfg["spatial_pool_size"])
120
+ self.cnn_proj = nn.Linear(self.cnn.feat_dim, hid)
121
+ self.pre_temporal_norm = nn.LayerNorm(hid)
122
+ self.temporal_blocks = nn.Sequential(
123
+ *[ResidualAttentionBlock(hid, cfg["num_heads"]) for _ in range(cfg["temporal_layers"])]
124
+ )
125
+ self.post_temporal = nn.Sequential(nn.ReLU(), nn.Linear(hid, hid), nn.LayerNorm(hid))
126
+ self.key_head = sub_head(hid, cfg["head_hidden"], cfg["dropout"], cfg["num_keys"])
127
+ assert cfg["rotation_head_type"] == "hybrid"
128
+ self.rotation_head = HybridRotationHead(hid, cfg["head_hidden"], cfg["dropout"])
129
+
130
+ def forward(self, x: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
131
+ """x: (B, N, 3, H, W) normalised frames -> key logits (B, N, 5), rotation raw (B, N, 26)."""
132
+ b, n = x.shape[:2]
133
+ f = self.cnn(x.flatten(0, 1)).view(b, n, -1)
134
+ f = self.post_temporal(self.temporal_blocks(self.pre_temporal_norm(self.cnn_proj(f))))
135
+ flat = f.reshape(b * n, -1)
136
+ return self.key_head(flat).view(b, n, -1), self.rotation_head(flat).view(b, n, -1)
137
+
138
+
139
+ def decode_hybrid_rotation(raw: np.ndarray) -> tuple[np.ndarray, np.ndarray]:
140
+ """(N, 26) -> (yaw_deg, pitch_deg): argmax bin centre plus residual."""
141
+ ny = len(YAW_BIN_CENTERS)
142
+ yaw = np.array(YAW_BIN_CENTERS, dtype=np.float32)[raw[:, :ny].argmax(1)] + raw[:, ny]
143
+ pitch_logits = raw[:, ny + 1 : ny + 1 + len(PITCH_BIN_CENTERS)]
144
+ pitch = np.array(PITCH_BIN_CENTERS, dtype=np.float32)[pitch_logits.argmax(1)] + raw[:, -1]
145
+ return yaw, pitch
146
+
147
+
148
+ def load_model(weights: Path, cfg: dict) -> IDMImpalaVPT:
149
+ model = IDMImpalaVPT(cfg)
150
+ model.load_state_dict(load_file(str(weights)), strict=True)
151
+ return model.eval()
152
+
153
+
154
+ def decode_frames(video: Path, resolution: int) -> tuple[np.ndarray, float]:
155
+ """All frames as (N, res, res, 3) uint8 RGB, squashed to a square with INTER_AREA."""
156
+ frames: list[np.ndarray] = []
157
+ with av.open(str(video)) as container:
158
+ rate = container.streams.video[0].average_rate
159
+ fps = float(rate) if rate else 24.0
160
+ for frame in container.decode(video=0):
161
+ img = frame.to_ndarray(format="rgb24")
162
+ frames.append(cv2.resize(img, (resolution, resolution), interpolation=cv2.INTER_AREA))
163
+ assert frames, f"no frames decoded from {video}"
164
+ return np.stack(frames), fps
165
+
166
+
167
+ def generate_windows(total: int, size: int, stride: int, min_frames: int) -> list[tuple[int, int]]:
168
+ if total < min_frames:
169
+ return []
170
+ if total < size:
171
+ return [(0, total)]
172
+ windows, start = [], 0
173
+ while start < total:
174
+ end = min(start + size, total)
175
+ windows.append((start, end))
176
+ if end == total:
177
+ break
178
+ start += stride
179
+ return windows
180
+
181
+
182
+ @torch.no_grad()
183
+ def predict(video: Path, weights: Path, cfg: dict, device: torch.device) -> dict:
184
+ model = load_model(weights, cfg).to(device)
185
+ frames, fps = decode_frames(video, cfg["resolution"])
186
+ total = len(frames)
187
+ windows = generate_windows(total, cfg["window_frames"], cfg["window_stride"], cfg["min_frames"])
188
+ assert windows, f"video has {total} frames; need at least {cfg['min_frames']}"
189
+ norm = torch.from_numpy(((frames.astype(np.float32) / 255.0 - IMAGENET_MEAN) / IMAGENET_STD)).permute(0, 3, 1, 2)
190
+ dim = cfg["num_keys"] + 26
191
+ acc, count = np.zeros((total, dim), np.float64), np.zeros(total, np.float64)
192
+ for start, end in windows:
193
+ x = norm[start:end].unsqueeze(0).to(device)
194
+ key_logits, rot = model(x)
195
+ raw = torch.cat([torch.sigmoid(key_logits), rot], dim=-1)[0].cpu().numpy()
196
+ acc[start:end] += raw
197
+ count[start:end] += 1
198
+ acc /= count[:, None]
199
+ nk = cfg["num_keys"]
200
+ yaw, pitch = decode_hybrid_rotation(acc[:, nk:].astype(np.float32))
201
+ prob = acc[:, :nk]
202
+ rows = [
203
+ {
204
+ "frame": i,
205
+ "t_s": round(i / fps, 3),
206
+ "key_probabilities": dict(zip(cfg["key_order"], (round(float(p), 4) for p in prob[i]))),
207
+ "keys_pressed": [k for k, p in zip(cfg["key_order"], prob[i]) if p > cfg["key_threshold"]],
208
+ "yaw_deg": round(float(yaw[i]), 4),
209
+ "pitch_deg": round(float(pitch[i]), 4),
210
+ }
211
+ for i in range(total)
212
+ ]
213
+ return {"video": str(video), "fps": fps, "num_frames": total, "frames": rows}
214
+
215
+
216
+ def main() -> None:
217
+ here = Path(__file__).parent
218
+ ap = argparse.ArgumentParser(description=__doc__)
219
+ ap.add_argument("video", type=Path)
220
+ ap.add_argument("--weights", type=Path, default=here / "model.safetensors")
221
+ ap.add_argument("--config", type=Path, default=here / "config.json")
222
+ ap.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu")
223
+ args = ap.parse_args()
224
+ cfg = json.loads(args.config.read_text())
225
+ print(json.dumps(predict(args.video, args.weights, cfg, torch.device(args.device)), indent=2))
226
+
227
+
228
+ if __name__ == "__main__":
229
+ main()
pixel/model.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:4498b70cb763b3595599a717d6a80a22b581356fee3a62dacf402aba38636953
3
+ size 39353324