liangsu9988 commited on
Commit
ee7544c
·
verified ·
1 Parent(s): 1aa712e

Promote latest kernel artifacts to main

Browse files
Files changed (23) hide show
  1. README.md +0 -9
  2. benchmarks/benchmark_native_parity.py +520 -0
  3. build/torch211-cxx11-cu128-x86_64-linux/__init__.py +374 -1
  4. build/torch211-cxx11-cu128-x86_64-linux/_ops.py +3 -3
  5. build/torch211-cxx11-cu128-x86_64-linux/{_world_model_conv_cuda_f14c443.abi3.so → _world_model_conv_cuda_33f8494.abi3.so} +2 -2
  6. build/torch211-cxx11-cu128-x86_64-linux/metadata.json +1 -1
  7. build/torch211-cxx11-cu130-aarch64-linux/__init__.py +430 -0
  8. build/torch211-cxx11-cu130-aarch64-linux/_ops.py +6 -0
  9. build/{torch212-cxx11-cu132-x86_64-linux/_world_model_conv_cuda_f14c443.abi3.so → torch211-cxx11-cu130-aarch64-linux/_world_model_conv_cuda_7781728.abi3.so} +2 -2
  10. build/torch211-cxx11-cu130-aarch64-linux/metadata.json +32 -0
  11. build/torch211-cxx11-cu130-aarch64-linux/world_model_conv/__init__.py +14 -0
  12. build/torch211-cxx11-cu130-x86_64-linux/__init__.py +374 -1
  13. build/torch211-cxx11-cu130-x86_64-linux/_ops.py +3 -3
  14. build/torch211-cxx11-cu130-x86_64-linux/{_world_model_conv_cuda_f14c443.abi3.so → _world_model_conv_cuda_33f8494.abi3.so} +2 -2
  15. build/torch211-cxx11-cu130-x86_64-linux/metadata.json +2 -1
  16. build/torch212-cxx11-cu130-x86_64-linux/__init__.py +374 -1
  17. build/torch212-cxx11-cu130-x86_64-linux/_ops.py +3 -3
  18. build/torch212-cxx11-cu130-x86_64-linux/{_world_model_conv_cuda_f14c443.abi3.so → _world_model_conv_cuda_33f8494.abi3.so} +2 -2
  19. build/torch212-cxx11-cu130-x86_64-linux/metadata.json +2 -1
  20. build/torch212-cxx11-cu132-x86_64-linux/__init__.py +374 -1
  21. build/torch212-cxx11-cu132-x86_64-linux/_ops.py +3 -3
  22. build/torch212-cxx11-cu132-x86_64-linux/_world_model_conv_cuda_33f8494.abi3.so +3 -0
  23. build/torch212-cxx11-cu132-x86_64-linux/metadata.json +2 -1
README.md DELETED
@@ -1,9 +0,0 @@
1
- # flashrt/world-model-conv
2
-
3
- This repository is a compatibility mirror for older `kernels` clients
4
- that resolve repositories through the default Hugging Face model repo API.
5
-
6
- Canonical Kernel Hub repo: https://huggingface.co/kernels/flashrt/world-model-conv
7
-
8
- Do not edit this mirror by hand. It is generated from the Kernel Hub
9
- `vN` branches and contains the same `build/**` artifacts.
 
 
 
 
 
 
 
 
 
 
benchmarks/benchmark_native_parity.py ADDED
@@ -0,0 +1,520 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ """World-model Conv benchmark with native, wrapper, compile, and cuDNN paths."""
3
+
4
+ from __future__ import annotations
5
+
6
+ import argparse
7
+ import json
8
+ import os
9
+ import sys
10
+ from dataclasses import asdict, dataclass
11
+ from pathlib import Path
12
+
13
+ import torch
14
+ import torch.nn.functional as F
15
+
16
+
17
+ PACKAGE = Path(__file__).resolve().parents[1]
18
+ sys.path.insert(0, str(PACKAGE / "tests"))
19
+ from test_world_model_conv import ( # noqa: E402
20
+ dequantize_linear_nvfp4,
21
+ load_installed_ops,
22
+ load_source_ops,
23
+ quantize_conv_tensor,
24
+ )
25
+
26
+
27
+ CONV3D_SHAPES = {
28
+ "causal-c32": (1, 2, 4, 16, 16, 32, 32),
29
+ "causal-small": (1, 2, 4, 16, 16, 64, 64),
30
+ }
31
+ NVFP4_CONV3D_SHAPES = {
32
+ "nvfp4-c64": (1, 2, 4, 16, 16, 64, 64),
33
+ "nvfp4-c128": (1, 2, 4, 16, 16, 128, 128),
34
+ "nvfp4-c512": (1, 2, 4, 16, 16, 512, 512),
35
+ }
36
+ CONV2D_SHAPES = {
37
+ "resample-c64": (4, 32, 32, 64, 64),
38
+ "resample-c320": (17, 32, 32, 320, 320),
39
+ }
40
+
41
+
42
+ @dataclass
43
+ class Result:
44
+ workload: str
45
+ shape: str
46
+ native_us: float
47
+ wrapper_us: float
48
+ wrapper_native: float
49
+ eager_cudnn_us: float
50
+ compile_cudnn_us: float
51
+ diagnostic_predequant_cudnn_us: float | None
52
+ diagnostic_predequant_compile_us: float | None
53
+ max_abs: float
54
+ mean_abs: float
55
+ p99_abs: float
56
+ cosine: float
57
+ accepted: bool
58
+
59
+
60
+ def bench(fn, warmup, iters):
61
+ for _ in range(warmup):
62
+ fn()
63
+ torch.cuda.synchronize()
64
+ start = torch.cuda.Event(enable_timing=True)
65
+ end = torch.cuda.Event(enable_timing=True)
66
+ start.record()
67
+ for _ in range(iters):
68
+ fn()
69
+ end.record()
70
+ torch.cuda.synchronize()
71
+ return start.elapsed_time(end) * 1000.0 / iters
72
+
73
+
74
+ def build_native():
75
+ from torch.utils.cpp_extension import load
76
+
77
+ os.environ.setdefault("TORCH_CUDA_ARCH_LIST", "12.0a")
78
+ return load(
79
+ name="world_model_conv_raw_native",
80
+ sources=[
81
+ str(PACKAGE / "benchmarks/native_binding.cpp"),
82
+ str(PACKAGE / "csrc/fp8_conv3d_sm120_v18.cu"),
83
+ str(PACKAGE / "csrc/fp8_causal_conv3d_sm120.cu"),
84
+ str(PACKAGE / "csrc/fp8_conv2d_3x3_sm120.cu"),
85
+ str(PACKAGE / "csrc/nvfp4_causal_conv3d_sm120.cu"),
86
+ str(PACKAGE / "csrc/nvfp4_causal_conv3d_residual_sm120.cu"),
87
+ str(PACKAGE / "csrc/nvfp4_causal_conv3d_residual_k128_sm120.cu"),
88
+ ],
89
+ extra_include_paths=[str(PACKAGE / "csrc")],
90
+ extra_cflags=["-O3"],
91
+ extra_cuda_cflags=["-O3"],
92
+ verbose=False,
93
+ )
94
+
95
+
96
+ def metrics(got, ref):
97
+ diff = (got.float() - ref.float()).abs().flatten()
98
+ cosine = F.cosine_similarity(
99
+ got.float().flatten(), ref.float().flatten(), dim=0
100
+ ).item()
101
+ return (
102
+ diff.max().item(),
103
+ diff.mean().item(),
104
+ torch.quantile(diff, 0.99).item(),
105
+ cosine,
106
+ )
107
+
108
+
109
+ def source_call(ops, name, *args, out):
110
+ if hasattr(ops, "_ops"):
111
+ getattr(ops._ops, name)(*args, out)
112
+ else:
113
+ getattr(ops, name)(*args, out=out)
114
+
115
+
116
+ def run_conv3d(ops, native, label, shape, args):
117
+ n, tc, tn, h, w, ci, co = shape
118
+ cache = (torch.randn((n, tc, h, w, ci), device="cuda") * 0.1).to(
119
+ torch.float8_e4m3fn
120
+ )
121
+ new = (torch.randn((n, tn, h, w, ci), device="cuda") * 0.1).to(
122
+ torch.float8_e4m3fn
123
+ )
124
+ weight = (torch.randn((co, 3, 3, 3, ci), device="cuda") * 0.1).to(
125
+ torch.float8_e4m3fn
126
+ )
127
+ bias = (torch.randn(co, device="cuda") * 0.01).to(torch.bfloat16)
128
+ out = torch.empty((n, tn, h, w, co), device="cuda", dtype=torch.bfloat16)
129
+ alpha = 0.75
130
+
131
+ wrapper = lambda: source_call(
132
+ ops,
133
+ "fp8_causal_conv3d_ndhwc_bf16",
134
+ cache,
135
+ new,
136
+ weight,
137
+ bias,
138
+ alpha,
139
+ out=out,
140
+ )
141
+ raw = lambda: native.causal_conv3d(
142
+ cache, new, weight, bias, alpha, out
143
+ )
144
+
145
+ def cudnn_ref():
146
+ x = torch.cat((cache, new), dim=1).float().permute(0, 4, 1, 2, 3)
147
+ wt = weight.float().permute(0, 4, 1, 2, 3)
148
+ y = F.conv3d(x, wt, padding=(0, 1, 1))
149
+ y = y.mul(alpha).add(bias.float().view(1, -1, 1, 1, 1))
150
+ out.copy_(y[:, :, :tn].permute(0, 2, 3, 4, 1).to(torch.bfloat16))
151
+
152
+ compiled = torch.compile(cudnn_ref, fullgraph=True)
153
+ wrapper()
154
+ got = out.clone()
155
+ cudnn_ref()
156
+ ref = out.clone()
157
+ max_abs, mean_abs, p99_abs, cosine = metrics(got, ref)
158
+ native_us = bench(raw, args.warmup, args.iters)
159
+ wrapper_us = bench(wrapper, args.warmup, args.iters)
160
+ eager_us = bench(cudnn_ref, args.warmup, args.iters)
161
+ compile_us = bench(compiled, args.warmup, args.iters)
162
+ return Result(
163
+ label,
164
+ str(shape),
165
+ native_us,
166
+ wrapper_us,
167
+ wrapper_us / native_us,
168
+ eager_us,
169
+ compile_us,
170
+ None,
171
+ None,
172
+ max_abs,
173
+ mean_abs,
174
+ p99_abs,
175
+ cosine,
176
+ wrapper_us - native_us <= max(0.5, native_us * 0.05)
177
+ and wrapper_us <= min(eager_us, compile_us) * 0.98
178
+ and cosine >= 0.999
179
+ and mean_abs <= 0.01,
180
+ )
181
+
182
+
183
+ def run_conv3d_residual(ops, native, label, shape, args):
184
+ n, tc, tn, h, w, ci, co = shape
185
+ cache = (torch.randn((n, tc, h, w, ci), device="cuda") * 0.1).to(
186
+ torch.float8_e4m3fn
187
+ )
188
+ new = (torch.randn((n, tn, h, w, ci), device="cuda") * 0.1).to(
189
+ torch.float8_e4m3fn
190
+ )
191
+ weight = (torch.randn((co, 3, 3, 3, ci), device="cuda") * 0.1).to(
192
+ torch.float8_e4m3fn
193
+ )
194
+ bias = (torch.randn(co, device="cuda") * 0.01).to(torch.bfloat16)
195
+ residual = torch.randn(
196
+ (n, co, tn, h, w), device="cuda", dtype=torch.bfloat16
197
+ )
198
+ out = torch.empty_like(residual)
199
+ alpha = 0.75
200
+ wrapper = lambda: source_call(
201
+ ops,
202
+ "fp8_conv3d_v18_ncdhw_res_bf16out",
203
+ cache,
204
+ new,
205
+ weight,
206
+ bias,
207
+ residual,
208
+ alpha,
209
+ out=out,
210
+ )
211
+ raw = lambda: native.causal_conv3d_residual(
212
+ cache, new, weight, bias, residual, alpha, out
213
+ )
214
+
215
+ def cudnn_ref():
216
+ x = torch.cat((cache, new), dim=1).float().permute(0, 4, 1, 2, 3)
217
+ wt = weight.float().permute(0, 4, 1, 2, 3)
218
+ y = F.conv3d(x, wt, padding=(0, 1, 1))
219
+ y = y.mul(alpha).add(bias.float().view(1, -1, 1, 1, 1))
220
+ y = (
221
+ y[:, :, :tn].to(torch.bfloat16).float()
222
+ + residual.float()
223
+ ).to(torch.bfloat16)
224
+ out.copy_(y)
225
+
226
+ compiled = torch.compile(cudnn_ref, fullgraph=True)
227
+ wrapper()
228
+ got = out.clone()
229
+ cudnn_ref()
230
+ ref = out.clone()
231
+ max_abs, mean_abs, p99_abs, cosine = metrics(got, ref)
232
+ native_us = bench(raw, args.warmup, args.iters)
233
+ wrapper_us = bench(wrapper, args.warmup, args.iters)
234
+ eager_us = bench(cudnn_ref, args.warmup, args.iters)
235
+ compile_us = bench(compiled, args.warmup, args.iters)
236
+ return Result(
237
+ f"{label}-residual",
238
+ str(shape),
239
+ native_us,
240
+ wrapper_us,
241
+ wrapper_us / native_us,
242
+ eager_us,
243
+ compile_us,
244
+ None,
245
+ None,
246
+ max_abs,
247
+ mean_abs,
248
+ p99_abs,
249
+ cosine,
250
+ wrapper_us - native_us <= max(0.5, native_us * 0.05)
251
+ and wrapper_us <= min(eager_us, compile_us) * 0.98
252
+ and cosine >= 0.999
253
+ and mean_abs <= 0.01,
254
+ )
255
+
256
+
257
+ def run_conv2d(ops, native, label, shape, args):
258
+ n, h, w, ci, co = shape
259
+ input = (torch.randn((n, h, w, ci), device="cuda") * 0.1).to(
260
+ torch.float8_e4m3fn
261
+ )
262
+ weight = (torch.randn((co, 3, 3, ci), device="cuda") * 0.1).to(
263
+ torch.float8_e4m3fn
264
+ )
265
+ bias = (torch.randn(co, device="cuda") * 0.01).to(torch.bfloat16)
266
+ out = torch.empty((n, h, w, co), device="cuda", dtype=torch.bfloat16)
267
+ alpha = 0.75
268
+ wrapper = lambda: source_call(
269
+ ops,
270
+ "fp8_conv2d_3x3_nhwc_bf16",
271
+ input,
272
+ weight,
273
+ bias,
274
+ alpha,
275
+ out=out,
276
+ )
277
+ raw = lambda: native.conv2d(input, weight, bias, alpha, out)
278
+
279
+ def cudnn_ref():
280
+ x = input.float().permute(0, 3, 1, 2)
281
+ wt = weight.float().permute(0, 3, 1, 2)
282
+ y = F.conv2d(x, wt, padding=1).mul(alpha)
283
+ y = y.add(bias.float().view(1, -1, 1, 1))
284
+ out.copy_(y.permute(0, 2, 3, 1).to(torch.bfloat16))
285
+
286
+ compiled = torch.compile(cudnn_ref, fullgraph=True)
287
+ wrapper()
288
+ got = out.clone()
289
+ cudnn_ref()
290
+ ref = out.clone()
291
+ max_abs, mean_abs, p99_abs, cosine = metrics(got, ref)
292
+ native_us = bench(raw, args.warmup, args.iters)
293
+ wrapper_us = bench(wrapper, args.warmup, args.iters)
294
+ eager_us = bench(cudnn_ref, args.warmup, args.iters)
295
+ compile_us = bench(compiled, args.warmup, args.iters)
296
+ return Result(
297
+ label,
298
+ str(shape),
299
+ native_us,
300
+ wrapper_us,
301
+ wrapper_us / native_us,
302
+ eager_us,
303
+ compile_us,
304
+ None,
305
+ None,
306
+ max_abs,
307
+ mean_abs,
308
+ p99_abs,
309
+ cosine,
310
+ wrapper_us - native_us <= max(0.5, native_us * 0.05)
311
+ and wrapper_us <= min(eager_us, compile_us) * 0.98
312
+ and cosine >= 0.999
313
+ and mean_abs <= 0.01,
314
+ )
315
+
316
+
317
+ def run_nvfp4_conv3d(ops, native, label, shape, args, residual_path):
318
+ n, tc, tn, h, w, ci, co = shape
319
+ cache_bf16 = (
320
+ torch.randn((n, tc, h, w, ci), device="cuda") * 0.1
321
+ ).to(torch.bfloat16)
322
+ input_bf16 = (
323
+ torch.randn((n, tn, h, w, ci), device="cuda") * 0.1
324
+ ).to(torch.bfloat16)
325
+ weight_bf16 = (
326
+ torch.randn((co, 3, 3, 3, ci), device="cuda") * 0.1
327
+ ).to(torch.bfloat16)
328
+ cache, cache_sf = quantize_conv_tensor(cache_bf16)
329
+ input, input_sf = quantize_conv_tensor(input_bf16)
330
+ weight, weight_sf = quantize_conv_tensor(weight_bf16)
331
+ cache_dequant = dequantize_linear_nvfp4(
332
+ cache.reshape(-1, ci // 2), cache_sf.reshape(-1, ci // 16)
333
+ ).reshape_as(cache_bf16)
334
+ input_dequant = dequantize_linear_nvfp4(
335
+ input.reshape(-1, ci // 2), input_sf.reshape(-1, ci // 16)
336
+ ).reshape_as(input_bf16)
337
+ weight_dequant = dequantize_linear_nvfp4(
338
+ weight.reshape(-1, ci // 2), weight_sf.reshape(-1, ci // 16)
339
+ ).reshape_as(weight_bf16)
340
+ bias = (torch.randn(co, device="cuda") * 0.01).to(torch.bfloat16)
341
+ alpha = 0.75
342
+
343
+ if residual_path:
344
+ residual = torch.randn(
345
+ (n, co, tn, h, w), device="cuda", dtype=torch.bfloat16
346
+ )
347
+ out = torch.empty_like(residual)
348
+ wrapper = lambda: source_call(
349
+ ops,
350
+ "nvfp4_causal_conv3d_residual_ncdhw_bf16",
351
+ cache, input, weight, cache_sf, input_sf, weight_sf, bias,
352
+ residual, None, alpha, out=out,
353
+ )
354
+ raw = lambda: native.nvfp4_causal_conv3d_residual(
355
+ cache, input, weight, cache_sf, input_sf, weight_sf, bias,
356
+ residual, alpha, out,
357
+ )
358
+ else:
359
+ residual = None
360
+ out = torch.empty(
361
+ (n, tn, h, w, co), device="cuda", dtype=torch.bfloat16
362
+ )
363
+ wrapper = lambda: source_call(
364
+ ops,
365
+ "nvfp4_causal_conv3d_ndhwc_bf16",
366
+ cache, input, weight, cache_sf, input_sf, weight_sf, bias,
367
+ None, alpha, out=out,
368
+ )
369
+ raw = lambda: native.nvfp4_causal_conv3d(
370
+ cache, input, weight, cache_sf, input_sf, weight_sf, bias,
371
+ alpha, out,
372
+ )
373
+
374
+ def store_cudnn_result(cache_value, input_value):
375
+ x = torch.cat((cache_value, input_value), dim=1).permute(
376
+ 0, 4, 1, 2, 3
377
+ )
378
+ wt = weight_dequant.permute(0, 4, 1, 2, 3)
379
+ y = F.conv3d(x, wt, padding=(0, 1, 1))[:, :, :tn]
380
+ y = y.mul(alpha).add(bias.float().view(1, -1, 1, 1, 1))
381
+ if residual_path:
382
+ out.copy_(
383
+ (y.to(torch.bfloat16).float() + residual.float()).to(
384
+ torch.bfloat16
385
+ )
386
+ )
387
+ else:
388
+ out.copy_(y.permute(0, 2, 3, 4, 1).to(torch.bfloat16))
389
+
390
+ def predequant_cudnn_ref():
391
+ store_cudnn_result(cache_dequant, input_dequant)
392
+
393
+ magnitude = torch.tensor(
394
+ [0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0], device="cuda"
395
+ )
396
+ scale_values, scale_bytes = [], []
397
+ for byte in list(range(0x78)) + [0xFE]:
398
+ exponent = (byte >> 3) & 0xF
399
+ mantissa = byte & 0x7
400
+ value = (
401
+ (mantissa / 8.0) * (2.0 ** -6)
402
+ if exponent == 0
403
+ else (1.0 + mantissa / 8.0) * (2.0 ** (exponent - 7))
404
+ )
405
+ scale_values.append(value)
406
+ scale_bytes.append(byte)
407
+ scale_lookup = torch.zeros(256, device="cuda")
408
+ scale_lookup[
409
+ torch.tensor(scale_bytes, device="cuda", dtype=torch.long)
410
+ ] = torch.tensor(scale_values, device="cuda")
411
+
412
+ def unpack(packed_value, scale_value):
413
+ low = packed_value & 0xF
414
+ high = packed_value >> 4
415
+ low_value = magnitude[(low & 0x7).long()] * torch.where(
416
+ low & 0x8 != 0, -1.0, 1.0
417
+ )
418
+ high_value = magnitude[(high & 0x7).long()] * torch.where(
419
+ high & 0x8 != 0, -1.0, 1.0
420
+ )
421
+ values = torch.stack((low_value, high_value), dim=-1).flatten(-2)
422
+ scales_value = scale_lookup[scale_value.long()].repeat_interleave(
423
+ 16, dim=-1
424
+ )
425
+ return values * scales_value
426
+
427
+ def cudnn_ref():
428
+ cache_value = unpack(cache, cache_sf).reshape_as(cache_bf16)
429
+ input_value = unpack(input, input_sf).reshape_as(input_bf16)
430
+ store_cudnn_result(cache_value, input_value)
431
+
432
+ compiled = torch.compile(cudnn_ref, fullgraph=True)
433
+ predequant_compiled = torch.compile(predequant_cudnn_ref, fullgraph=True)
434
+ wrapper()
435
+ got = out.clone()
436
+ cudnn_ref()
437
+ ref = out.clone()
438
+ max_abs, mean_abs, p99_abs, cosine = metrics(got, ref)
439
+ native_us = bench(raw, args.warmup, args.iters)
440
+ wrapper_us = bench(wrapper, args.warmup, args.iters)
441
+ eager_us = bench(cudnn_ref, args.warmup, args.iters)
442
+ compile_us = bench(compiled, args.warmup, args.iters)
443
+ diagnostic_eager_us = bench(
444
+ predequant_cudnn_ref, args.warmup, args.iters
445
+ )
446
+ diagnostic_compile_us = bench(
447
+ predequant_compiled, args.warmup, args.iters
448
+ )
449
+ return Result(
450
+ f"{label}{'-residual' if residual_path else ''}",
451
+ str(shape),
452
+ native_us,
453
+ wrapper_us,
454
+ wrapper_us / native_us,
455
+ eager_us,
456
+ compile_us,
457
+ diagnostic_eager_us,
458
+ diagnostic_compile_us,
459
+ max_abs,
460
+ mean_abs,
461
+ p99_abs,
462
+ cosine,
463
+ wrapper_us - native_us <= max(0.5, native_us * 0.05)
464
+ and wrapper_us <= min(eager_us, compile_us) * 0.98
465
+ and cosine >= 0.998
466
+ and mean_abs <= 0.02,
467
+ )
468
+
469
+
470
+ def main():
471
+ parser = argparse.ArgumentParser()
472
+ parser.add_argument("--backend", choices=["source", "installed"], default="source")
473
+ parser.add_argument("--artifact")
474
+ parser.add_argument("--warmup", type=int, default=10)
475
+ parser.add_argument("--iters", type=int, default=30)
476
+ parser.add_argument("--output")
477
+ args = parser.parse_args()
478
+ ops = (
479
+ load_source_ops()
480
+ if args.backend == "source"
481
+ else load_installed_ops(args.artifact)
482
+ )
483
+ native = build_native()
484
+ rows = [
485
+ *(run_conv3d(ops, native, name, shape, args)
486
+ for name, shape in CONV3D_SHAPES.items()),
487
+ *(run_conv3d_residual(ops, native, name, shape, args)
488
+ for name, shape in CONV3D_SHAPES.items()
489
+ if shape[-1] % 8 == 0),
490
+ *(run_conv2d(ops, native, name, shape, args)
491
+ for name, shape in CONV2D_SHAPES.items()),
492
+ *(run_nvfp4_conv3d(ops, native, name, shape, args, False)
493
+ for name, shape in NVFP4_CONV3D_SHAPES.items()),
494
+ *(run_nvfp4_conv3d(ops, native, name, shape, args, True)
495
+ for name, shape in NVFP4_CONV3D_SHAPES.items()),
496
+ ]
497
+ for row in rows:
498
+ print(
499
+ f"{row.workload}: native={row.native_us:.3f}us "
500
+ f"wrapper={row.wrapper_us:.3f}us ({row.wrapper_native:.3f}) "
501
+ f"cuDNN-eager={row.eager_cudnn_us:.3f}us "
502
+ f"cuDNN-compile={row.compile_cudnn_us:.3f}us "
503
+ + (
504
+ f"predequant-compile="
505
+ f"{row.diagnostic_predequant_compile_us:.3f}us "
506
+ if row.diagnostic_predequant_compile_us is not None
507
+ else ""
508
+ )
509
+ + f"cos={row.cosine:.7f} accepted={row.accepted}"
510
+ )
511
+ if args.output:
512
+ path = Path(args.output)
513
+ path.parent.mkdir(parents=True, exist_ok=True)
514
+ path.write_text(json.dumps([asdict(row) for row in rows], indent=2) + "\n")
515
+ if not all(row.accepted for row in rows):
516
+ raise SystemExit("world-model Conv acceptance failed")
517
+
518
+
519
+ if __name__ == "__main__":
520
+ main()
build/torch211-cxx11-cu128-x86_64-linux/__init__.py CHANGED
@@ -9,6 +9,21 @@ import torch
9
  from ._ops import add_op_namespace_prefix, ops
10
 
11
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
12
  @torch.library.register_fake(add_op_namespace_prefix("fp8_conv3d_v18_ncdhw_res_bf16out"))
13
  def _fp8_conv3d_fake(
14
  cache_x: torch.Tensor,
@@ -27,6 +42,11 @@ def _fp8_conv3d_fake(
27
  raise RuntimeError("cache_x must have shape (N,2,H,W,Ci)")
28
  if weight.shape != (co, 3, 3, 3, ci):
29
  raise RuntimeError("weight must have shape (Co,3,3,3,Ci)")
 
 
 
 
 
30
  if residual.shape != (n, co, t_new, h, w) or out.shape != residual.shape:
31
  raise RuntimeError("residual/out must be NCDHW")
32
  if bias.shape != (co,):
@@ -34,6 +54,179 @@ def _fp8_conv3d_fake(
34
  return None
35
 
36
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
37
  def fp8_conv3d_v18_ncdhw_res_bf16out(
38
  cache_x: torch.Tensor,
39
  new_x: torch.Tensor,
@@ -54,4 +247,184 @@ def fp8_conv3d_v18_ncdhw_res_bf16out(
54
  return out
55
 
56
 
57
- __all__ = ["fp8_conv3d_v18_ncdhw_res_bf16out"]
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
9
  from ._ops import add_op_namespace_prefix, ops
10
 
11
 
12
+ @torch.library.register_fake(
13
+ add_op_namespace_prefix("bf16_causal_conv3d_ndhwc_bf16")
14
+ )
15
+ def _bf16_causal_conv3d_fake(cache_x, new_x, weight, bias, alpha, out) -> None:
16
+ n, t, h, w, ci = new_x.shape
17
+ co = weight.shape[0]
18
+ if cache_x.shape != (n, 2, h, w, ci):
19
+ raise RuntimeError("cache_x must have shape (N,2,H,W,Ci)")
20
+ if weight.shape != (co, 3, 3, 3, ci):
21
+ raise RuntimeError("weight must have shape (Co,3,3,3,Ci)")
22
+ if bias.shape != (co,) or out.shape != (n, t, h, w, co):
23
+ raise RuntimeError("bias/out shape mismatch")
24
+ return None
25
+
26
+
27
  @torch.library.register_fake(add_op_namespace_prefix("fp8_conv3d_v18_ncdhw_res_bf16out"))
28
  def _fp8_conv3d_fake(
29
  cache_x: torch.Tensor,
 
42
  raise RuntimeError("cache_x must have shape (N,2,H,W,Ci)")
43
  if weight.shape != (co, 3, 3, 3, ci):
44
  raise RuntimeError("weight must have shape (Co,3,3,3,Ci)")
45
+ if ci not in (32, 64):
46
+ raise RuntimeError(
47
+ "FP8 Conv3D is accepted only for Ci=32/64; use the "
48
+ "strong-library or NVFP4 path for larger channels"
49
+ )
50
  if residual.shape != (n, co, t_new, h, w) or out.shape != residual.shape:
51
  raise RuntimeError("residual/out must be NCDHW")
52
  if bias.shape != (co,):
 
54
  return None
55
 
56
 
57
+ @torch.library.register_fake(
58
+ add_op_namespace_prefix("fp8_causal_conv3d_ndhwc_bf16")
59
+ )
60
+ def _fp8_causal_conv3d_fake(
61
+ cache_x: torch.Tensor,
62
+ new_x: torch.Tensor,
63
+ weight: torch.Tensor,
64
+ bias: torch.Tensor,
65
+ alpha: float,
66
+ out: torch.Tensor,
67
+ ) -> None:
68
+ if cache_x.dim() != 5 or new_x.dim() != 5:
69
+ raise RuntimeError("cache_x/new_x must be NDHWC")
70
+ n, t, h, w, ci = new_x.shape
71
+ co = weight.shape[0]
72
+ if cache_x.shape != (n, 2, h, w, ci):
73
+ raise RuntimeError("cache_x must have shape (N,2,H,W,Ci)")
74
+ if weight.shape != (co, 3, 3, 3, ci):
75
+ raise RuntimeError("weight must have shape (Co,3,3,3,Ci)")
76
+ if ci not in (32, 64):
77
+ raise RuntimeError(
78
+ "FP8 Conv3D is accepted only for Ci=32/64; use the "
79
+ "strong-library or NVFP4 path for larger channels"
80
+ )
81
+ if bias.shape != (co,) or out.shape != (n, t, h, w, co):
82
+ raise RuntimeError("bias/out shape mismatch")
83
+ return None
84
+
85
+
86
+ @torch.library.register_fake(
87
+ add_op_namespace_prefix("fp8_conv2d_3x3_nhwc_bf16")
88
+ )
89
+ def _fp8_conv2d_fake(
90
+ input: torch.Tensor,
91
+ weight: torch.Tensor,
92
+ bias: torch.Tensor,
93
+ alpha: float,
94
+ out: torch.Tensor,
95
+ ) -> None:
96
+ if input.dim() != 4:
97
+ raise RuntimeError("input must have shape (N,H,W,Ci)")
98
+ n, h, w, ci = input.shape
99
+ co = weight.shape[0]
100
+ if weight.shape != (co, 3, 3, ci):
101
+ raise RuntimeError("weight must have shape (Co,3,3,Ci)")
102
+ if bias.shape != (co,) or out.shape != (n, h, w, co):
103
+ raise RuntimeError("bias/out shape mismatch")
104
+ return None
105
+
106
+
107
+ @torch.library.register_fake(
108
+ add_op_namespace_prefix("fp8_conv2d_3x3_ncdhw_bf16")
109
+ )
110
+ def _fp8_conv2d_ncdhw_fake(
111
+ input: torch.Tensor,
112
+ weight: torch.Tensor,
113
+ bias: torch.Tensor,
114
+ alpha: float,
115
+ out: torch.Tensor,
116
+ ) -> None:
117
+ if input.dim() != 5:
118
+ raise RuntimeError("input must have shape (B,T,H,W,Ci)")
119
+ b, t, h, w, ci = input.shape
120
+ co = weight.shape[0]
121
+ if weight.shape != (co, 3, 3, ci):
122
+ raise RuntimeError("weight must have shape (Co,3,3,Ci)")
123
+ if bias.shape != (co,) or out.shape != (b, co, t, h, w):
124
+ raise RuntimeError("bias/out shape mismatch")
125
+ return None
126
+
127
+
128
+ def _check_nvfp4_conv_shapes(
129
+ cache_packed: torch.Tensor,
130
+ new_packed: torch.Tensor,
131
+ weight_packed: torch.Tensor,
132
+ cache_sf: torch.Tensor,
133
+ new_sf: torch.Tensor,
134
+ weight_sf: torch.Tensor,
135
+ bias: torch.Tensor,
136
+ outer_weight: torch.Tensor | None,
137
+ ) -> tuple[int, int, int, int, int]:
138
+ if (
139
+ cache_packed.dim() != 5
140
+ or new_packed.dim() != 5
141
+ or weight_packed.dim() != 5
142
+ ):
143
+ raise RuntimeError("packed inputs must be five-dimensional")
144
+ n, t, h, w, ci_half = new_packed.shape
145
+ ci = ci_half * 2
146
+ co = weight_packed.shape[0]
147
+ if (
148
+ (ci != 64 and ci % 128 != 0)
149
+ or co % 8 != 0
150
+ or cache_packed.shape != (n, 2, h, w, ci // 2)
151
+ or weight_packed.shape != (co, 3, 3, 3, ci // 2)
152
+ or cache_sf.shape != (n, 2, h, w, ci // 16)
153
+ or new_sf.shape != (n, t, h, w, ci // 16)
154
+ or weight_sf.shape != (co, 3, 3, 3, ci // 16)
155
+ or bias.shape != (co,)
156
+ or (outer_weight is not None and outer_weight.shape != (co,))
157
+ ):
158
+ raise RuntimeError(
159
+ "NVFP4 Conv3D accepts Ci=64 or multiples of 128; "
160
+ "other shapes must use the strong-library path"
161
+ )
162
+ return n, t, h, w, co
163
+
164
+
165
+ @torch.library.register_fake(
166
+ add_op_namespace_prefix("nvfp4_causal_conv3d_ndhwc_bf16")
167
+ )
168
+ def _nvfp4_causal_conv3d_fake(
169
+ cache_packed: torch.Tensor,
170
+ new_packed: torch.Tensor,
171
+ weight_packed: torch.Tensor,
172
+ cache_sf: torch.Tensor,
173
+ new_sf: torch.Tensor,
174
+ weight_sf: torch.Tensor,
175
+ bias: torch.Tensor,
176
+ outer_weight: torch.Tensor | None,
177
+ alpha: float,
178
+ out: torch.Tensor,
179
+ ) -> None:
180
+ del alpha
181
+ n, t, h, w, co = _check_nvfp4_conv_shapes(
182
+ cache_packed,
183
+ new_packed,
184
+ weight_packed,
185
+ cache_sf,
186
+ new_sf,
187
+ weight_sf,
188
+ bias,
189
+ outer_weight,
190
+ )
191
+ if out.shape != (n, t, h, w, co):
192
+ raise RuntimeError("out must have shape (N,T,H,W,Co)")
193
+ return None
194
+
195
+
196
+ @torch.library.register_fake(
197
+ add_op_namespace_prefix(
198
+ "nvfp4_causal_conv3d_residual_ncdhw_bf16"
199
+ )
200
+ )
201
+ def _nvfp4_causal_conv3d_residual_fake(
202
+ cache_packed: torch.Tensor,
203
+ new_packed: torch.Tensor,
204
+ weight_packed: torch.Tensor,
205
+ cache_sf: torch.Tensor,
206
+ new_sf: torch.Tensor,
207
+ weight_sf: torch.Tensor,
208
+ bias: torch.Tensor,
209
+ residual: torch.Tensor,
210
+ outer_weight: torch.Tensor | None,
211
+ alpha: float,
212
+ out: torch.Tensor,
213
+ ) -> None:
214
+ del alpha
215
+ n, t, h, w, co = _check_nvfp4_conv_shapes(
216
+ cache_packed,
217
+ new_packed,
218
+ weight_packed,
219
+ cache_sf,
220
+ new_sf,
221
+ weight_sf,
222
+ bias,
223
+ outer_weight,
224
+ )
225
+ if residual.shape != (n, co, t, h, w) or out.shape != residual.shape:
226
+ raise RuntimeError("residual/out must have shape (N,Co,T,H,W)")
227
+ return None
228
+
229
+
230
  def fp8_conv3d_v18_ncdhw_res_bf16out(
231
  cache_x: torch.Tensor,
232
  new_x: torch.Tensor,
 
247
  return out
248
 
249
 
250
+ def bf16_causal_conv3d_ndhwc_bf16(
251
+ cache_x: torch.Tensor,
252
+ new_x: torch.Tensor,
253
+ weight: torch.Tensor,
254
+ bias: torch.Tensor,
255
+ alpha: float = 1.0,
256
+ *,
257
+ out: Optional[torch.Tensor] = None,
258
+ ) -> torch.Tensor:
259
+ """Experimental SM110 BF16 causal Conv3D probe.
260
+
261
+ This entry is explicit opt-in. It is not the default world-model Conv3D
262
+ backend because the current native kernel does not beat cuDNN.
263
+ """
264
+ n, t, h, w, _ = new_x.shape
265
+ if out is None:
266
+ out = torch.empty(
267
+ (n, t, h, w, weight.shape[0]),
268
+ device=new_x.device,
269
+ dtype=torch.bfloat16,
270
+ )
271
+ ops.bf16_causal_conv3d_ndhwc_bf16(
272
+ cache_x, new_x, weight, bias, float(alpha), out
273
+ )
274
+ return out
275
+
276
+
277
+ def fp8_causal_conv3d_ndhwc_bf16(
278
+ cache_x: torch.Tensor,
279
+ new_x: torch.Tensor,
280
+ weight: torch.Tensor,
281
+ bias: torch.Tensor,
282
+ alpha: float = 1.0,
283
+ *,
284
+ out: Optional[torch.Tensor] = None,
285
+ ) -> torch.Tensor:
286
+ """FP8 causal 3D convolution with virtual two-frame cache concat."""
287
+
288
+ n, t, h, w, _ = new_x.shape
289
+ if out is None:
290
+ out = torch.empty(
291
+ (n, t, h, w, weight.shape[0]),
292
+ device=new_x.device,
293
+ dtype=torch.bfloat16,
294
+ )
295
+ ops.fp8_causal_conv3d_ndhwc_bf16(
296
+ cache_x, new_x, weight, bias, float(alpha), out
297
+ )
298
+ return out
299
+
300
+
301
+ def fp8_conv2d_3x3_nhwc_bf16(
302
+ input: torch.Tensor,
303
+ weight: torch.Tensor,
304
+ bias: torch.Tensor,
305
+ alpha: float = 1.0,
306
+ *,
307
+ out: Optional[torch.Tensor] = None,
308
+ ) -> torch.Tensor:
309
+ """FP8 3x3 Conv2D with NHWC input/output and BF16 epilogue."""
310
+
311
+ if out is None:
312
+ out = torch.empty(
313
+ (*input.shape[:3], weight.shape[0]),
314
+ device=input.device,
315
+ dtype=torch.bfloat16,
316
+ )
317
+ ops.fp8_conv2d_3x3_nhwc_bf16(
318
+ input, weight, bias, float(alpha), out
319
+ )
320
+ return out
321
+
322
+
323
+ def fp8_conv2d_3x3_ncdhw_bf16(
324
+ input: torch.Tensor,
325
+ weight: torch.Tensor,
326
+ bias: torch.Tensor,
327
+ alpha: float = 1.0,
328
+ *,
329
+ out: Optional[torch.Tensor] = None,
330
+ ) -> torch.Tensor:
331
+ """FP8 3x3 Conv2D over B*T frames with direct BF16 NCDHW output."""
332
+
333
+ if out is None:
334
+ out = torch.empty(
335
+ (
336
+ input.shape[0],
337
+ weight.shape[0],
338
+ input.shape[1],
339
+ input.shape[2],
340
+ input.shape[3],
341
+ ),
342
+ device=input.device,
343
+ dtype=torch.bfloat16,
344
+ )
345
+ ops.fp8_conv2d_3x3_ncdhw_bf16(
346
+ input, weight, bias, float(alpha), out
347
+ )
348
+ return out
349
+
350
+
351
+ def nvfp4_causal_conv3d_ndhwc_bf16(
352
+ cache_packed: torch.Tensor,
353
+ new_packed: torch.Tensor,
354
+ weight_packed: torch.Tensor,
355
+ cache_sf: torch.Tensor,
356
+ new_sf: torch.Tensor,
357
+ weight_sf: torch.Tensor,
358
+ bias: torch.Tensor,
359
+ outer_weight: torch.Tensor | None = None,
360
+ alpha: float = 1.0,
361
+ *,
362
+ out: Optional[torch.Tensor] = None,
363
+ ) -> torch.Tensor:
364
+ """NVFP4 causal Conv3D with linear UE4M3 scale factors and BF16 NDHWC output."""
365
+
366
+ n, t, h, w, _ = new_packed.shape
367
+ if out is None:
368
+ out = torch.empty(
369
+ (n, t, h, w, weight_packed.shape[0]),
370
+ device=new_packed.device,
371
+ dtype=torch.bfloat16,
372
+ )
373
+ ops.nvfp4_causal_conv3d_ndhwc_bf16(
374
+ cache_packed,
375
+ new_packed,
376
+ weight_packed,
377
+ cache_sf,
378
+ new_sf,
379
+ weight_sf,
380
+ bias,
381
+ outer_weight,
382
+ float(alpha),
383
+ out,
384
+ )
385
+ return out
386
+
387
+
388
+ def nvfp4_causal_conv3d_residual_ncdhw_bf16(
389
+ cache_packed: torch.Tensor,
390
+ new_packed: torch.Tensor,
391
+ weight_packed: torch.Tensor,
392
+ cache_sf: torch.Tensor,
393
+ new_sf: torch.Tensor,
394
+ weight_sf: torch.Tensor,
395
+ bias: torch.Tensor,
396
+ residual: torch.Tensor,
397
+ outer_weight: torch.Tensor | None = None,
398
+ alpha: float = 1.0,
399
+ *,
400
+ out: Optional[torch.Tensor] = None,
401
+ ) -> torch.Tensor:
402
+ """NVFP4 causal Conv3D with fused bias, residual, and BF16 NCDHW output."""
403
+
404
+ if out is None:
405
+ out = torch.empty_like(residual)
406
+ ops.nvfp4_causal_conv3d_residual_ncdhw_bf16(
407
+ cache_packed,
408
+ new_packed,
409
+ weight_packed,
410
+ cache_sf,
411
+ new_sf,
412
+ weight_sf,
413
+ bias,
414
+ residual,
415
+ outer_weight,
416
+ float(alpha),
417
+ out,
418
+ )
419
+ return out
420
+
421
+
422
+ __all__ = [
423
+ "bf16_causal_conv3d_ndhwc_bf16",
424
+ "fp8_conv3d_v18_ncdhw_res_bf16out",
425
+ "fp8_causal_conv3d_ndhwc_bf16",
426
+ "fp8_conv2d_3x3_nhwc_bf16",
427
+ "fp8_conv2d_3x3_ncdhw_bf16",
428
+ "nvfp4_causal_conv3d_ndhwc_bf16",
429
+ "nvfp4_causal_conv3d_residual_ncdhw_bf16",
430
+ ]
build/torch211-cxx11-cu128-x86_64-linux/_ops.py CHANGED
@@ -1,9 +1,9 @@
1
  import torch
2
- from . import _world_model_conv_cuda_f14c443
3
- ops = torch.ops._world_model_conv_cuda_f14c443
4
 
5
  def add_op_namespace_prefix(op_name: str):
6
  """
7
  Prefix op by namespace.
8
  """
9
- return f"_world_model_conv_cuda_f14c443::{op_name}"
 
1
  import torch
2
+ from . import _world_model_conv_cuda_33f8494
3
+ ops = torch.ops._world_model_conv_cuda_33f8494
4
 
5
  def add_op_namespace_prefix(op_name: str):
6
  """
7
  Prefix op by namespace.
8
  """
9
+ return f"_world_model_conv_cuda_33f8494::{op_name}"
build/torch211-cxx11-cu128-x86_64-linux/{_world_model_conv_cuda_f14c443.abi3.so → _world_model_conv_cuda_33f8494.abi3.so} RENAMED
@@ -1,3 +1,3 @@
1
  version https://git-lfs.github.com/spec/v1
2
- oid sha256:724ae0ca8e8d4638c2d8aae226836993124662348462c1d83e6f558109bdf221
3
- size 296248
 
1
  version https://git-lfs.github.com/spec/v1
2
+ oid sha256:56c0b4484f40ef55e8dbf7a4e6ebc872c6f9dc5985e510419c945ee52d624d88
3
+ size 1280336
build/torch211-cxx11-cu128-x86_64-linux/metadata.json CHANGED
@@ -1,6 +1,6 @@
1
  {
2
  "name": "world-model-conv",
3
- "id": "_world_model_conv_cuda_f14c443",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "python-depends": [],
 
1
  {
2
  "name": "world-model-conv",
3
+ "id": "_world_model_conv_cuda_33f8494",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "python-depends": [],
build/torch211-cxx11-cu130-aarch64-linux/__init__.py ADDED
@@ -0,0 +1,430 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """FlashRT world-model convolution kernels."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from typing import Optional
6
+
7
+ import torch
8
+
9
+ from ._ops import add_op_namespace_prefix, ops
10
+
11
+
12
+ @torch.library.register_fake(
13
+ add_op_namespace_prefix("bf16_causal_conv3d_ndhwc_bf16")
14
+ )
15
+ def _bf16_causal_conv3d_fake(cache_x, new_x, weight, bias, alpha, out) -> None:
16
+ n, t, h, w, ci = new_x.shape
17
+ co = weight.shape[0]
18
+ if cache_x.shape != (n, 2, h, w, ci):
19
+ raise RuntimeError("cache_x must have shape (N,2,H,W,Ci)")
20
+ if weight.shape != (co, 3, 3, 3, ci):
21
+ raise RuntimeError("weight must have shape (Co,3,3,3,Ci)")
22
+ if bias.shape != (co,) or out.shape != (n, t, h, w, co):
23
+ raise RuntimeError("bias/out shape mismatch")
24
+ return None
25
+
26
+
27
+ @torch.library.register_fake(add_op_namespace_prefix("fp8_conv3d_v18_ncdhw_res_bf16out"))
28
+ def _fp8_conv3d_fake(
29
+ cache_x: torch.Tensor,
30
+ new_x: torch.Tensor,
31
+ weight: torch.Tensor,
32
+ bias: torch.Tensor,
33
+ residual: torch.Tensor,
34
+ alpha: float,
35
+ out: torch.Tensor,
36
+ ) -> None:
37
+ if cache_x.dim() != 5 or new_x.dim() != 5:
38
+ raise RuntimeError("cache_x/new_x must be NDHWC")
39
+ n, t_new, h, w, ci = new_x.shape
40
+ co = weight.shape[0]
41
+ if cache_x.shape != (n, 2, h, w, ci):
42
+ raise RuntimeError("cache_x must have shape (N,2,H,W,Ci)")
43
+ if weight.shape != (co, 3, 3, 3, ci):
44
+ raise RuntimeError("weight must have shape (Co,3,3,3,Ci)")
45
+ if ci not in (32, 64):
46
+ raise RuntimeError(
47
+ "FP8 Conv3D is accepted only for Ci=32/64; use the "
48
+ "strong-library or NVFP4 path for larger channels"
49
+ )
50
+ if residual.shape != (n, co, t_new, h, w) or out.shape != residual.shape:
51
+ raise RuntimeError("residual/out must be NCDHW")
52
+ if bias.shape != (co,):
53
+ raise RuntimeError("bias must have shape (Co,)")
54
+ return None
55
+
56
+
57
+ @torch.library.register_fake(
58
+ add_op_namespace_prefix("fp8_causal_conv3d_ndhwc_bf16")
59
+ )
60
+ def _fp8_causal_conv3d_fake(
61
+ cache_x: torch.Tensor,
62
+ new_x: torch.Tensor,
63
+ weight: torch.Tensor,
64
+ bias: torch.Tensor,
65
+ alpha: float,
66
+ out: torch.Tensor,
67
+ ) -> None:
68
+ if cache_x.dim() != 5 or new_x.dim() != 5:
69
+ raise RuntimeError("cache_x/new_x must be NDHWC")
70
+ n, t, h, w, ci = new_x.shape
71
+ co = weight.shape[0]
72
+ if cache_x.shape != (n, 2, h, w, ci):
73
+ raise RuntimeError("cache_x must have shape (N,2,H,W,Ci)")
74
+ if weight.shape != (co, 3, 3, 3, ci):
75
+ raise RuntimeError("weight must have shape (Co,3,3,3,Ci)")
76
+ if ci not in (32, 64):
77
+ raise RuntimeError(
78
+ "FP8 Conv3D is accepted only for Ci=32/64; use the "
79
+ "strong-library or NVFP4 path for larger channels"
80
+ )
81
+ if bias.shape != (co,) or out.shape != (n, t, h, w, co):
82
+ raise RuntimeError("bias/out shape mismatch")
83
+ return None
84
+
85
+
86
+ @torch.library.register_fake(
87
+ add_op_namespace_prefix("fp8_conv2d_3x3_nhwc_bf16")
88
+ )
89
+ def _fp8_conv2d_fake(
90
+ input: torch.Tensor,
91
+ weight: torch.Tensor,
92
+ bias: torch.Tensor,
93
+ alpha: float,
94
+ out: torch.Tensor,
95
+ ) -> None:
96
+ if input.dim() != 4:
97
+ raise RuntimeError("input must have shape (N,H,W,Ci)")
98
+ n, h, w, ci = input.shape
99
+ co = weight.shape[0]
100
+ if weight.shape != (co, 3, 3, ci):
101
+ raise RuntimeError("weight must have shape (Co,3,3,Ci)")
102
+ if bias.shape != (co,) or out.shape != (n, h, w, co):
103
+ raise RuntimeError("bias/out shape mismatch")
104
+ return None
105
+
106
+
107
+ @torch.library.register_fake(
108
+ add_op_namespace_prefix("fp8_conv2d_3x3_ncdhw_bf16")
109
+ )
110
+ def _fp8_conv2d_ncdhw_fake(
111
+ input: torch.Tensor,
112
+ weight: torch.Tensor,
113
+ bias: torch.Tensor,
114
+ alpha: float,
115
+ out: torch.Tensor,
116
+ ) -> None:
117
+ if input.dim() != 5:
118
+ raise RuntimeError("input must have shape (B,T,H,W,Ci)")
119
+ b, t, h, w, ci = input.shape
120
+ co = weight.shape[0]
121
+ if weight.shape != (co, 3, 3, ci):
122
+ raise RuntimeError("weight must have shape (Co,3,3,Ci)")
123
+ if bias.shape != (co,) or out.shape != (b, co, t, h, w):
124
+ raise RuntimeError("bias/out shape mismatch")
125
+ return None
126
+
127
+
128
+ def _check_nvfp4_conv_shapes(
129
+ cache_packed: torch.Tensor,
130
+ new_packed: torch.Tensor,
131
+ weight_packed: torch.Tensor,
132
+ cache_sf: torch.Tensor,
133
+ new_sf: torch.Tensor,
134
+ weight_sf: torch.Tensor,
135
+ bias: torch.Tensor,
136
+ outer_weight: torch.Tensor | None,
137
+ ) -> tuple[int, int, int, int, int]:
138
+ if (
139
+ cache_packed.dim() != 5
140
+ or new_packed.dim() != 5
141
+ or weight_packed.dim() != 5
142
+ ):
143
+ raise RuntimeError("packed inputs must be five-dimensional")
144
+ n, t, h, w, ci_half = new_packed.shape
145
+ ci = ci_half * 2
146
+ co = weight_packed.shape[0]
147
+ if (
148
+ (ci != 64 and ci % 128 != 0)
149
+ or co % 8 != 0
150
+ or cache_packed.shape != (n, 2, h, w, ci // 2)
151
+ or weight_packed.shape != (co, 3, 3, 3, ci // 2)
152
+ or cache_sf.shape != (n, 2, h, w, ci // 16)
153
+ or new_sf.shape != (n, t, h, w, ci // 16)
154
+ or weight_sf.shape != (co, 3, 3, 3, ci // 16)
155
+ or bias.shape != (co,)
156
+ or (outer_weight is not None and outer_weight.shape != (co,))
157
+ ):
158
+ raise RuntimeError(
159
+ "NVFP4 Conv3D accepts Ci=64 or multiples of 128; "
160
+ "other shapes must use the strong-library path"
161
+ )
162
+ return n, t, h, w, co
163
+
164
+
165
+ @torch.library.register_fake(
166
+ add_op_namespace_prefix("nvfp4_causal_conv3d_ndhwc_bf16")
167
+ )
168
+ def _nvfp4_causal_conv3d_fake(
169
+ cache_packed: torch.Tensor,
170
+ new_packed: torch.Tensor,
171
+ weight_packed: torch.Tensor,
172
+ cache_sf: torch.Tensor,
173
+ new_sf: torch.Tensor,
174
+ weight_sf: torch.Tensor,
175
+ bias: torch.Tensor,
176
+ outer_weight: torch.Tensor | None,
177
+ alpha: float,
178
+ out: torch.Tensor,
179
+ ) -> None:
180
+ del alpha
181
+ n, t, h, w, co = _check_nvfp4_conv_shapes(
182
+ cache_packed,
183
+ new_packed,
184
+ weight_packed,
185
+ cache_sf,
186
+ new_sf,
187
+ weight_sf,
188
+ bias,
189
+ outer_weight,
190
+ )
191
+ if out.shape != (n, t, h, w, co):
192
+ raise RuntimeError("out must have shape (N,T,H,W,Co)")
193
+ return None
194
+
195
+
196
+ @torch.library.register_fake(
197
+ add_op_namespace_prefix(
198
+ "nvfp4_causal_conv3d_residual_ncdhw_bf16"
199
+ )
200
+ )
201
+ def _nvfp4_causal_conv3d_residual_fake(
202
+ cache_packed: torch.Tensor,
203
+ new_packed: torch.Tensor,
204
+ weight_packed: torch.Tensor,
205
+ cache_sf: torch.Tensor,
206
+ new_sf: torch.Tensor,
207
+ weight_sf: torch.Tensor,
208
+ bias: torch.Tensor,
209
+ residual: torch.Tensor,
210
+ outer_weight: torch.Tensor | None,
211
+ alpha: float,
212
+ out: torch.Tensor,
213
+ ) -> None:
214
+ del alpha
215
+ n, t, h, w, co = _check_nvfp4_conv_shapes(
216
+ cache_packed,
217
+ new_packed,
218
+ weight_packed,
219
+ cache_sf,
220
+ new_sf,
221
+ weight_sf,
222
+ bias,
223
+ outer_weight,
224
+ )
225
+ if residual.shape != (n, co, t, h, w) or out.shape != residual.shape:
226
+ raise RuntimeError("residual/out must have shape (N,Co,T,H,W)")
227
+ return None
228
+
229
+
230
+ def fp8_conv3d_v18_ncdhw_res_bf16out(
231
+ cache_x: torch.Tensor,
232
+ new_x: torch.Tensor,
233
+ weight: torch.Tensor,
234
+ bias: torch.Tensor,
235
+ residual: torch.Tensor,
236
+ alpha: float = 1.0,
237
+ *,
238
+ out: Optional[torch.Tensor] = None,
239
+ ) -> torch.Tensor:
240
+ """FP8 3D causal conv with virtual cache concat, bias, residual, BF16 NCDHW output."""
241
+
242
+ n, t_new, h, w, _ = new_x.shape
243
+ co = weight.shape[0]
244
+ if out is None:
245
+ out = torch.empty((n, co, t_new, h, w), device=new_x.device, dtype=torch.bfloat16)
246
+ ops.fp8_conv3d_v18_ncdhw_res_bf16out(cache_x, new_x, weight, bias, residual, float(alpha), out)
247
+ return out
248
+
249
+
250
+ def bf16_causal_conv3d_ndhwc_bf16(
251
+ cache_x: torch.Tensor,
252
+ new_x: torch.Tensor,
253
+ weight: torch.Tensor,
254
+ bias: torch.Tensor,
255
+ alpha: float = 1.0,
256
+ *,
257
+ out: Optional[torch.Tensor] = None,
258
+ ) -> torch.Tensor:
259
+ """Experimental SM110 BF16 causal Conv3D probe.
260
+
261
+ This entry is explicit opt-in. It is not the default world-model Conv3D
262
+ backend because the current native kernel does not beat cuDNN.
263
+ """
264
+ n, t, h, w, _ = new_x.shape
265
+ if out is None:
266
+ out = torch.empty(
267
+ (n, t, h, w, weight.shape[0]),
268
+ device=new_x.device,
269
+ dtype=torch.bfloat16,
270
+ )
271
+ ops.bf16_causal_conv3d_ndhwc_bf16(
272
+ cache_x, new_x, weight, bias, float(alpha), out
273
+ )
274
+ return out
275
+
276
+
277
+ def fp8_causal_conv3d_ndhwc_bf16(
278
+ cache_x: torch.Tensor,
279
+ new_x: torch.Tensor,
280
+ weight: torch.Tensor,
281
+ bias: torch.Tensor,
282
+ alpha: float = 1.0,
283
+ *,
284
+ out: Optional[torch.Tensor] = None,
285
+ ) -> torch.Tensor:
286
+ """FP8 causal 3D convolution with virtual two-frame cache concat."""
287
+
288
+ n, t, h, w, _ = new_x.shape
289
+ if out is None:
290
+ out = torch.empty(
291
+ (n, t, h, w, weight.shape[0]),
292
+ device=new_x.device,
293
+ dtype=torch.bfloat16,
294
+ )
295
+ ops.fp8_causal_conv3d_ndhwc_bf16(
296
+ cache_x, new_x, weight, bias, float(alpha), out
297
+ )
298
+ return out
299
+
300
+
301
+ def fp8_conv2d_3x3_nhwc_bf16(
302
+ input: torch.Tensor,
303
+ weight: torch.Tensor,
304
+ bias: torch.Tensor,
305
+ alpha: float = 1.0,
306
+ *,
307
+ out: Optional[torch.Tensor] = None,
308
+ ) -> torch.Tensor:
309
+ """FP8 3x3 Conv2D with NHWC input/output and BF16 epilogue."""
310
+
311
+ if out is None:
312
+ out = torch.empty(
313
+ (*input.shape[:3], weight.shape[0]),
314
+ device=input.device,
315
+ dtype=torch.bfloat16,
316
+ )
317
+ ops.fp8_conv2d_3x3_nhwc_bf16(
318
+ input, weight, bias, float(alpha), out
319
+ )
320
+ return out
321
+
322
+
323
+ def fp8_conv2d_3x3_ncdhw_bf16(
324
+ input: torch.Tensor,
325
+ weight: torch.Tensor,
326
+ bias: torch.Tensor,
327
+ alpha: float = 1.0,
328
+ *,
329
+ out: Optional[torch.Tensor] = None,
330
+ ) -> torch.Tensor:
331
+ """FP8 3x3 Conv2D over B*T frames with direct BF16 NCDHW output."""
332
+
333
+ if out is None:
334
+ out = torch.empty(
335
+ (
336
+ input.shape[0],
337
+ weight.shape[0],
338
+ input.shape[1],
339
+ input.shape[2],
340
+ input.shape[3],
341
+ ),
342
+ device=input.device,
343
+ dtype=torch.bfloat16,
344
+ )
345
+ ops.fp8_conv2d_3x3_ncdhw_bf16(
346
+ input, weight, bias, float(alpha), out
347
+ )
348
+ return out
349
+
350
+
351
+ def nvfp4_causal_conv3d_ndhwc_bf16(
352
+ cache_packed: torch.Tensor,
353
+ new_packed: torch.Tensor,
354
+ weight_packed: torch.Tensor,
355
+ cache_sf: torch.Tensor,
356
+ new_sf: torch.Tensor,
357
+ weight_sf: torch.Tensor,
358
+ bias: torch.Tensor,
359
+ outer_weight: torch.Tensor | None = None,
360
+ alpha: float = 1.0,
361
+ *,
362
+ out: Optional[torch.Tensor] = None,
363
+ ) -> torch.Tensor:
364
+ """NVFP4 causal Conv3D with linear UE4M3 scale factors and BF16 NDHWC output."""
365
+
366
+ n, t, h, w, _ = new_packed.shape
367
+ if out is None:
368
+ out = torch.empty(
369
+ (n, t, h, w, weight_packed.shape[0]),
370
+ device=new_packed.device,
371
+ dtype=torch.bfloat16,
372
+ )
373
+ ops.nvfp4_causal_conv3d_ndhwc_bf16(
374
+ cache_packed,
375
+ new_packed,
376
+ weight_packed,
377
+ cache_sf,
378
+ new_sf,
379
+ weight_sf,
380
+ bias,
381
+ outer_weight,
382
+ float(alpha),
383
+ out,
384
+ )
385
+ return out
386
+
387
+
388
+ def nvfp4_causal_conv3d_residual_ncdhw_bf16(
389
+ cache_packed: torch.Tensor,
390
+ new_packed: torch.Tensor,
391
+ weight_packed: torch.Tensor,
392
+ cache_sf: torch.Tensor,
393
+ new_sf: torch.Tensor,
394
+ weight_sf: torch.Tensor,
395
+ bias: torch.Tensor,
396
+ residual: torch.Tensor,
397
+ outer_weight: torch.Tensor | None = None,
398
+ alpha: float = 1.0,
399
+ *,
400
+ out: Optional[torch.Tensor] = None,
401
+ ) -> torch.Tensor:
402
+ """NVFP4 causal Conv3D with fused bias, residual, and BF16 NCDHW output."""
403
+
404
+ if out is None:
405
+ out = torch.empty_like(residual)
406
+ ops.nvfp4_causal_conv3d_residual_ncdhw_bf16(
407
+ cache_packed,
408
+ new_packed,
409
+ weight_packed,
410
+ cache_sf,
411
+ new_sf,
412
+ weight_sf,
413
+ bias,
414
+ residual,
415
+ outer_weight,
416
+ float(alpha),
417
+ out,
418
+ )
419
+ return out
420
+
421
+
422
+ __all__ = [
423
+ "bf16_causal_conv3d_ndhwc_bf16",
424
+ "fp8_conv3d_v18_ncdhw_res_bf16out",
425
+ "fp8_causal_conv3d_ndhwc_bf16",
426
+ "fp8_conv2d_3x3_nhwc_bf16",
427
+ "fp8_conv2d_3x3_ncdhw_bf16",
428
+ "nvfp4_causal_conv3d_ndhwc_bf16",
429
+ "nvfp4_causal_conv3d_residual_ncdhw_bf16",
430
+ ]
build/torch211-cxx11-cu130-aarch64-linux/_ops.py ADDED
@@ -0,0 +1,6 @@
 
 
 
 
 
 
 
1
+ import torch
2
+ from . import _world_model_conv_cuda_7781728
3
+ ops = torch.ops._world_model_conv_cuda_7781728
4
+
5
+ def add_op_namespace_prefix(op_name: str):
6
+ return f"_world_model_conv_cuda_7781728::{op_name}"
build/{torch212-cxx11-cu132-x86_64-linux/_world_model_conv_cuda_f14c443.abi3.so → torch211-cxx11-cu130-aarch64-linux/_world_model_conv_cuda_7781728.abi3.so} RENAMED
@@ -1,3 +1,3 @@
1
  version https://git-lfs.github.com/spec/v1
2
- oid sha256:e6b9e2cf2cc2247919f98ed62ca2977c7d90d1ded187bed45718a67de1ac581e
3
- size 304192
 
1
  version https://git-lfs.github.com/spec/v1
2
+ oid sha256:1304037cec073cba13e0a29977c2ee81f10faa05061f31aec5f4511365fdb967
3
+ size 248600
build/torch211-cxx11-cu130-aarch64-linux/metadata.json ADDED
@@ -0,0 +1,32 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "name": "world-model-conv",
3
+ "id": "_world_model_conv_cuda_7781728",
4
+ "version": 1,
5
+ "license": "Apache-2.0",
6
+ "python-depends": [],
7
+ "backend": {
8
+ "type": "cuda",
9
+ "archs": [
10
+ "11.0a"
11
+ ]
12
+ },
13
+ "digest": {
14
+ "algorithm": "sha256",
15
+ "files": {
16
+ "__init__.py": "04caWU2syFAi2nTxD3Utrqn/ihEMR0PWAdr/Leh+wG8=",
17
+ "_world_model_conv_cuda_7781728.abi3.so": "EwQDfOwHPLoT4KKZd8LugfEPqgUGHzGuxfRRE2X9uWc=",
18
+ "_ops.py": "wmgfXGhBOcXkdpY5RfPHQtx9h5nzVzRU/dId0Fb7Qcw=",
19
+ "world_model_conv/__init__.py": "v6p5XMfQzddhi1fLSAw4HX9CyS0rQsidvu9VsT01xi4="
20
+ }
21
+ },
22
+ "provenance": {
23
+ "kernel": {
24
+ "sha": "33f8494",
25
+ "dirty": false
26
+ },
27
+ "validation": {
28
+ "torch": "2.11.0+cu130",
29
+ "cuda": "13.0"
30
+ }
31
+ }
32
+ }
build/torch211-cxx11-cu130-aarch64-linux/world_model_conv/__init__.py ADDED
@@ -0,0 +1,14 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import ctypes
2
+ import importlib.util
3
+ import sys
4
+ from pathlib import Path
5
+
6
+ def _import_from_path(file_path: Path):
7
+ path_hash = '{:x}'.format(ctypes.c_size_t(hash(file_path.absolute())).value)
8
+ spec = importlib.util.spec_from_file_location(path_hash, file_path)
9
+ module = importlib.util.module_from_spec(spec)
10
+ sys.modules[path_hash] = module
11
+ spec.loader.exec_module(module)
12
+ return module
13
+
14
+ globals().update(vars(_import_from_path(Path(__file__).parent.parent / '__init__.py')))
build/torch211-cxx11-cu130-x86_64-linux/__init__.py CHANGED
@@ -9,6 +9,21 @@ import torch
9
  from ._ops import add_op_namespace_prefix, ops
10
 
11
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
12
  @torch.library.register_fake(add_op_namespace_prefix("fp8_conv3d_v18_ncdhw_res_bf16out"))
13
  def _fp8_conv3d_fake(
14
  cache_x: torch.Tensor,
@@ -27,6 +42,11 @@ def _fp8_conv3d_fake(
27
  raise RuntimeError("cache_x must have shape (N,2,H,W,Ci)")
28
  if weight.shape != (co, 3, 3, 3, ci):
29
  raise RuntimeError("weight must have shape (Co,3,3,3,Ci)")
 
 
 
 
 
30
  if residual.shape != (n, co, t_new, h, w) or out.shape != residual.shape:
31
  raise RuntimeError("residual/out must be NCDHW")
32
  if bias.shape != (co,):
@@ -34,6 +54,179 @@ def _fp8_conv3d_fake(
34
  return None
35
 
36
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
37
  def fp8_conv3d_v18_ncdhw_res_bf16out(
38
  cache_x: torch.Tensor,
39
  new_x: torch.Tensor,
@@ -54,4 +247,184 @@ def fp8_conv3d_v18_ncdhw_res_bf16out(
54
  return out
55
 
56
 
57
- __all__ = ["fp8_conv3d_v18_ncdhw_res_bf16out"]
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
9
  from ._ops import add_op_namespace_prefix, ops
10
 
11
 
12
+ @torch.library.register_fake(
13
+ add_op_namespace_prefix("bf16_causal_conv3d_ndhwc_bf16")
14
+ )
15
+ def _bf16_causal_conv3d_fake(cache_x, new_x, weight, bias, alpha, out) -> None:
16
+ n, t, h, w, ci = new_x.shape
17
+ co = weight.shape[0]
18
+ if cache_x.shape != (n, 2, h, w, ci):
19
+ raise RuntimeError("cache_x must have shape (N,2,H,W,Ci)")
20
+ if weight.shape != (co, 3, 3, 3, ci):
21
+ raise RuntimeError("weight must have shape (Co,3,3,3,Ci)")
22
+ if bias.shape != (co,) or out.shape != (n, t, h, w, co):
23
+ raise RuntimeError("bias/out shape mismatch")
24
+ return None
25
+
26
+
27
  @torch.library.register_fake(add_op_namespace_prefix("fp8_conv3d_v18_ncdhw_res_bf16out"))
28
  def _fp8_conv3d_fake(
29
  cache_x: torch.Tensor,
 
42
  raise RuntimeError("cache_x must have shape (N,2,H,W,Ci)")
43
  if weight.shape != (co, 3, 3, 3, ci):
44
  raise RuntimeError("weight must have shape (Co,3,3,3,Ci)")
45
+ if ci not in (32, 64):
46
+ raise RuntimeError(
47
+ "FP8 Conv3D is accepted only for Ci=32/64; use the "
48
+ "strong-library or NVFP4 path for larger channels"
49
+ )
50
  if residual.shape != (n, co, t_new, h, w) or out.shape != residual.shape:
51
  raise RuntimeError("residual/out must be NCDHW")
52
  if bias.shape != (co,):
 
54
  return None
55
 
56
 
57
+ @torch.library.register_fake(
58
+ add_op_namespace_prefix("fp8_causal_conv3d_ndhwc_bf16")
59
+ )
60
+ def _fp8_causal_conv3d_fake(
61
+ cache_x: torch.Tensor,
62
+ new_x: torch.Tensor,
63
+ weight: torch.Tensor,
64
+ bias: torch.Tensor,
65
+ alpha: float,
66
+ out: torch.Tensor,
67
+ ) -> None:
68
+ if cache_x.dim() != 5 or new_x.dim() != 5:
69
+ raise RuntimeError("cache_x/new_x must be NDHWC")
70
+ n, t, h, w, ci = new_x.shape
71
+ co = weight.shape[0]
72
+ if cache_x.shape != (n, 2, h, w, ci):
73
+ raise RuntimeError("cache_x must have shape (N,2,H,W,Ci)")
74
+ if weight.shape != (co, 3, 3, 3, ci):
75
+ raise RuntimeError("weight must have shape (Co,3,3,3,Ci)")
76
+ if ci not in (32, 64):
77
+ raise RuntimeError(
78
+ "FP8 Conv3D is accepted only for Ci=32/64; use the "
79
+ "strong-library or NVFP4 path for larger channels"
80
+ )
81
+ if bias.shape != (co,) or out.shape != (n, t, h, w, co):
82
+ raise RuntimeError("bias/out shape mismatch")
83
+ return None
84
+
85
+
86
+ @torch.library.register_fake(
87
+ add_op_namespace_prefix("fp8_conv2d_3x3_nhwc_bf16")
88
+ )
89
+ def _fp8_conv2d_fake(
90
+ input: torch.Tensor,
91
+ weight: torch.Tensor,
92
+ bias: torch.Tensor,
93
+ alpha: float,
94
+ out: torch.Tensor,
95
+ ) -> None:
96
+ if input.dim() != 4:
97
+ raise RuntimeError("input must have shape (N,H,W,Ci)")
98
+ n, h, w, ci = input.shape
99
+ co = weight.shape[0]
100
+ if weight.shape != (co, 3, 3, ci):
101
+ raise RuntimeError("weight must have shape (Co,3,3,Ci)")
102
+ if bias.shape != (co,) or out.shape != (n, h, w, co):
103
+ raise RuntimeError("bias/out shape mismatch")
104
+ return None
105
+
106
+
107
+ @torch.library.register_fake(
108
+ add_op_namespace_prefix("fp8_conv2d_3x3_ncdhw_bf16")
109
+ )
110
+ def _fp8_conv2d_ncdhw_fake(
111
+ input: torch.Tensor,
112
+ weight: torch.Tensor,
113
+ bias: torch.Tensor,
114
+ alpha: float,
115
+ out: torch.Tensor,
116
+ ) -> None:
117
+ if input.dim() != 5:
118
+ raise RuntimeError("input must have shape (B,T,H,W,Ci)")
119
+ b, t, h, w, ci = input.shape
120
+ co = weight.shape[0]
121
+ if weight.shape != (co, 3, 3, ci):
122
+ raise RuntimeError("weight must have shape (Co,3,3,Ci)")
123
+ if bias.shape != (co,) or out.shape != (b, co, t, h, w):
124
+ raise RuntimeError("bias/out shape mismatch")
125
+ return None
126
+
127
+
128
+ def _check_nvfp4_conv_shapes(
129
+ cache_packed: torch.Tensor,
130
+ new_packed: torch.Tensor,
131
+ weight_packed: torch.Tensor,
132
+ cache_sf: torch.Tensor,
133
+ new_sf: torch.Tensor,
134
+ weight_sf: torch.Tensor,
135
+ bias: torch.Tensor,
136
+ outer_weight: torch.Tensor | None,
137
+ ) -> tuple[int, int, int, int, int]:
138
+ if (
139
+ cache_packed.dim() != 5
140
+ or new_packed.dim() != 5
141
+ or weight_packed.dim() != 5
142
+ ):
143
+ raise RuntimeError("packed inputs must be five-dimensional")
144
+ n, t, h, w, ci_half = new_packed.shape
145
+ ci = ci_half * 2
146
+ co = weight_packed.shape[0]
147
+ if (
148
+ (ci != 64 and ci % 128 != 0)
149
+ or co % 8 != 0
150
+ or cache_packed.shape != (n, 2, h, w, ci // 2)
151
+ or weight_packed.shape != (co, 3, 3, 3, ci // 2)
152
+ or cache_sf.shape != (n, 2, h, w, ci // 16)
153
+ or new_sf.shape != (n, t, h, w, ci // 16)
154
+ or weight_sf.shape != (co, 3, 3, 3, ci // 16)
155
+ or bias.shape != (co,)
156
+ or (outer_weight is not None and outer_weight.shape != (co,))
157
+ ):
158
+ raise RuntimeError(
159
+ "NVFP4 Conv3D accepts Ci=64 or multiples of 128; "
160
+ "other shapes must use the strong-library path"
161
+ )
162
+ return n, t, h, w, co
163
+
164
+
165
+ @torch.library.register_fake(
166
+ add_op_namespace_prefix("nvfp4_causal_conv3d_ndhwc_bf16")
167
+ )
168
+ def _nvfp4_causal_conv3d_fake(
169
+ cache_packed: torch.Tensor,
170
+ new_packed: torch.Tensor,
171
+ weight_packed: torch.Tensor,
172
+ cache_sf: torch.Tensor,
173
+ new_sf: torch.Tensor,
174
+ weight_sf: torch.Tensor,
175
+ bias: torch.Tensor,
176
+ outer_weight: torch.Tensor | None,
177
+ alpha: float,
178
+ out: torch.Tensor,
179
+ ) -> None:
180
+ del alpha
181
+ n, t, h, w, co = _check_nvfp4_conv_shapes(
182
+ cache_packed,
183
+ new_packed,
184
+ weight_packed,
185
+ cache_sf,
186
+ new_sf,
187
+ weight_sf,
188
+ bias,
189
+ outer_weight,
190
+ )
191
+ if out.shape != (n, t, h, w, co):
192
+ raise RuntimeError("out must have shape (N,T,H,W,Co)")
193
+ return None
194
+
195
+
196
+ @torch.library.register_fake(
197
+ add_op_namespace_prefix(
198
+ "nvfp4_causal_conv3d_residual_ncdhw_bf16"
199
+ )
200
+ )
201
+ def _nvfp4_causal_conv3d_residual_fake(
202
+ cache_packed: torch.Tensor,
203
+ new_packed: torch.Tensor,
204
+ weight_packed: torch.Tensor,
205
+ cache_sf: torch.Tensor,
206
+ new_sf: torch.Tensor,
207
+ weight_sf: torch.Tensor,
208
+ bias: torch.Tensor,
209
+ residual: torch.Tensor,
210
+ outer_weight: torch.Tensor | None,
211
+ alpha: float,
212
+ out: torch.Tensor,
213
+ ) -> None:
214
+ del alpha
215
+ n, t, h, w, co = _check_nvfp4_conv_shapes(
216
+ cache_packed,
217
+ new_packed,
218
+ weight_packed,
219
+ cache_sf,
220
+ new_sf,
221
+ weight_sf,
222
+ bias,
223
+ outer_weight,
224
+ )
225
+ if residual.shape != (n, co, t, h, w) or out.shape != residual.shape:
226
+ raise RuntimeError("residual/out must have shape (N,Co,T,H,W)")
227
+ return None
228
+
229
+
230
  def fp8_conv3d_v18_ncdhw_res_bf16out(
231
  cache_x: torch.Tensor,
232
  new_x: torch.Tensor,
 
247
  return out
248
 
249
 
250
+ def bf16_causal_conv3d_ndhwc_bf16(
251
+ cache_x: torch.Tensor,
252
+ new_x: torch.Tensor,
253
+ weight: torch.Tensor,
254
+ bias: torch.Tensor,
255
+ alpha: float = 1.0,
256
+ *,
257
+ out: Optional[torch.Tensor] = None,
258
+ ) -> torch.Tensor:
259
+ """Experimental SM110 BF16 causal Conv3D probe.
260
+
261
+ This entry is explicit opt-in. It is not the default world-model Conv3D
262
+ backend because the current native kernel does not beat cuDNN.
263
+ """
264
+ n, t, h, w, _ = new_x.shape
265
+ if out is None:
266
+ out = torch.empty(
267
+ (n, t, h, w, weight.shape[0]),
268
+ device=new_x.device,
269
+ dtype=torch.bfloat16,
270
+ )
271
+ ops.bf16_causal_conv3d_ndhwc_bf16(
272
+ cache_x, new_x, weight, bias, float(alpha), out
273
+ )
274
+ return out
275
+
276
+
277
+ def fp8_causal_conv3d_ndhwc_bf16(
278
+ cache_x: torch.Tensor,
279
+ new_x: torch.Tensor,
280
+ weight: torch.Tensor,
281
+ bias: torch.Tensor,
282
+ alpha: float = 1.0,
283
+ *,
284
+ out: Optional[torch.Tensor] = None,
285
+ ) -> torch.Tensor:
286
+ """FP8 causal 3D convolution with virtual two-frame cache concat."""
287
+
288
+ n, t, h, w, _ = new_x.shape
289
+ if out is None:
290
+ out = torch.empty(
291
+ (n, t, h, w, weight.shape[0]),
292
+ device=new_x.device,
293
+ dtype=torch.bfloat16,
294
+ )
295
+ ops.fp8_causal_conv3d_ndhwc_bf16(
296
+ cache_x, new_x, weight, bias, float(alpha), out
297
+ )
298
+ return out
299
+
300
+
301
+ def fp8_conv2d_3x3_nhwc_bf16(
302
+ input: torch.Tensor,
303
+ weight: torch.Tensor,
304
+ bias: torch.Tensor,
305
+ alpha: float = 1.0,
306
+ *,
307
+ out: Optional[torch.Tensor] = None,
308
+ ) -> torch.Tensor:
309
+ """FP8 3x3 Conv2D with NHWC input/output and BF16 epilogue."""
310
+
311
+ if out is None:
312
+ out = torch.empty(
313
+ (*input.shape[:3], weight.shape[0]),
314
+ device=input.device,
315
+ dtype=torch.bfloat16,
316
+ )
317
+ ops.fp8_conv2d_3x3_nhwc_bf16(
318
+ input, weight, bias, float(alpha), out
319
+ )
320
+ return out
321
+
322
+
323
+ def fp8_conv2d_3x3_ncdhw_bf16(
324
+ input: torch.Tensor,
325
+ weight: torch.Tensor,
326
+ bias: torch.Tensor,
327
+ alpha: float = 1.0,
328
+ *,
329
+ out: Optional[torch.Tensor] = None,
330
+ ) -> torch.Tensor:
331
+ """FP8 3x3 Conv2D over B*T frames with direct BF16 NCDHW output."""
332
+
333
+ if out is None:
334
+ out = torch.empty(
335
+ (
336
+ input.shape[0],
337
+ weight.shape[0],
338
+ input.shape[1],
339
+ input.shape[2],
340
+ input.shape[3],
341
+ ),
342
+ device=input.device,
343
+ dtype=torch.bfloat16,
344
+ )
345
+ ops.fp8_conv2d_3x3_ncdhw_bf16(
346
+ input, weight, bias, float(alpha), out
347
+ )
348
+ return out
349
+
350
+
351
+ def nvfp4_causal_conv3d_ndhwc_bf16(
352
+ cache_packed: torch.Tensor,
353
+ new_packed: torch.Tensor,
354
+ weight_packed: torch.Tensor,
355
+ cache_sf: torch.Tensor,
356
+ new_sf: torch.Tensor,
357
+ weight_sf: torch.Tensor,
358
+ bias: torch.Tensor,
359
+ outer_weight: torch.Tensor | None = None,
360
+ alpha: float = 1.0,
361
+ *,
362
+ out: Optional[torch.Tensor] = None,
363
+ ) -> torch.Tensor:
364
+ """NVFP4 causal Conv3D with linear UE4M3 scale factors and BF16 NDHWC output."""
365
+
366
+ n, t, h, w, _ = new_packed.shape
367
+ if out is None:
368
+ out = torch.empty(
369
+ (n, t, h, w, weight_packed.shape[0]),
370
+ device=new_packed.device,
371
+ dtype=torch.bfloat16,
372
+ )
373
+ ops.nvfp4_causal_conv3d_ndhwc_bf16(
374
+ cache_packed,
375
+ new_packed,
376
+ weight_packed,
377
+ cache_sf,
378
+ new_sf,
379
+ weight_sf,
380
+ bias,
381
+ outer_weight,
382
+ float(alpha),
383
+ out,
384
+ )
385
+ return out
386
+
387
+
388
+ def nvfp4_causal_conv3d_residual_ncdhw_bf16(
389
+ cache_packed: torch.Tensor,
390
+ new_packed: torch.Tensor,
391
+ weight_packed: torch.Tensor,
392
+ cache_sf: torch.Tensor,
393
+ new_sf: torch.Tensor,
394
+ weight_sf: torch.Tensor,
395
+ bias: torch.Tensor,
396
+ residual: torch.Tensor,
397
+ outer_weight: torch.Tensor | None = None,
398
+ alpha: float = 1.0,
399
+ *,
400
+ out: Optional[torch.Tensor] = None,
401
+ ) -> torch.Tensor:
402
+ """NVFP4 causal Conv3D with fused bias, residual, and BF16 NCDHW output."""
403
+
404
+ if out is None:
405
+ out = torch.empty_like(residual)
406
+ ops.nvfp4_causal_conv3d_residual_ncdhw_bf16(
407
+ cache_packed,
408
+ new_packed,
409
+ weight_packed,
410
+ cache_sf,
411
+ new_sf,
412
+ weight_sf,
413
+ bias,
414
+ residual,
415
+ outer_weight,
416
+ float(alpha),
417
+ out,
418
+ )
419
+ return out
420
+
421
+
422
+ __all__ = [
423
+ "bf16_causal_conv3d_ndhwc_bf16",
424
+ "fp8_conv3d_v18_ncdhw_res_bf16out",
425
+ "fp8_causal_conv3d_ndhwc_bf16",
426
+ "fp8_conv2d_3x3_nhwc_bf16",
427
+ "fp8_conv2d_3x3_ncdhw_bf16",
428
+ "nvfp4_causal_conv3d_ndhwc_bf16",
429
+ "nvfp4_causal_conv3d_residual_ncdhw_bf16",
430
+ ]
build/torch211-cxx11-cu130-x86_64-linux/_ops.py CHANGED
@@ -1,9 +1,9 @@
1
  import torch
2
- from . import _world_model_conv_cuda_f14c443
3
- ops = torch.ops._world_model_conv_cuda_f14c443
4
 
5
  def add_op_namespace_prefix(op_name: str):
6
  """
7
  Prefix op by namespace.
8
  """
9
- return f"_world_model_conv_cuda_f14c443::{op_name}"
 
1
  import torch
2
+ from . import _world_model_conv_cuda_33f8494
3
+ ops = torch.ops._world_model_conv_cuda_33f8494
4
 
5
  def add_op_namespace_prefix(op_name: str):
6
  """
7
  Prefix op by namespace.
8
  """
9
+ return f"_world_model_conv_cuda_33f8494::{op_name}"
build/torch211-cxx11-cu130-x86_64-linux/{_world_model_conv_cuda_f14c443.abi3.so → _world_model_conv_cuda_33f8494.abi3.so} RENAMED
@@ -1,3 +1,3 @@
1
  version https://git-lfs.github.com/spec/v1
2
- oid sha256:5bf4f3051e364b68f0fa536f7d9a7199a71c14a64f4863ab0485ef3cde7ffd7c
3
- size 285088
 
1
  version https://git-lfs.github.com/spec/v1
2
+ oid sha256:09a791c961123323294a35768fd135c97e2c2758e740aaf6b221f65b2977734f
3
+ size 1289672
build/torch211-cxx11-cu130-x86_64-linux/metadata.json CHANGED
@@ -1,12 +1,13 @@
1
  {
2
  "name": "world-model-conv",
3
- "id": "_world_model_conv_cuda_f14c443",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "python-depends": [],
7
  "backend": {
8
  "type": "cuda",
9
  "archs": [
 
10
  "12.0a"
11
  ]
12
  }
 
1
  {
2
  "name": "world-model-conv",
3
+ "id": "_world_model_conv_cuda_33f8494",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "python-depends": [],
7
  "backend": {
8
  "type": "cuda",
9
  "archs": [
10
+ "11.0a",
11
  "12.0a"
12
  ]
13
  }
build/torch212-cxx11-cu130-x86_64-linux/__init__.py CHANGED
@@ -9,6 +9,21 @@ import torch
9
  from ._ops import add_op_namespace_prefix, ops
10
 
11
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
12
  @torch.library.register_fake(add_op_namespace_prefix("fp8_conv3d_v18_ncdhw_res_bf16out"))
13
  def _fp8_conv3d_fake(
14
  cache_x: torch.Tensor,
@@ -27,6 +42,11 @@ def _fp8_conv3d_fake(
27
  raise RuntimeError("cache_x must have shape (N,2,H,W,Ci)")
28
  if weight.shape != (co, 3, 3, 3, ci):
29
  raise RuntimeError("weight must have shape (Co,3,3,3,Ci)")
 
 
 
 
 
30
  if residual.shape != (n, co, t_new, h, w) or out.shape != residual.shape:
31
  raise RuntimeError("residual/out must be NCDHW")
32
  if bias.shape != (co,):
@@ -34,6 +54,179 @@ def _fp8_conv3d_fake(
34
  return None
35
 
36
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
37
  def fp8_conv3d_v18_ncdhw_res_bf16out(
38
  cache_x: torch.Tensor,
39
  new_x: torch.Tensor,
@@ -54,4 +247,184 @@ def fp8_conv3d_v18_ncdhw_res_bf16out(
54
  return out
55
 
56
 
57
- __all__ = ["fp8_conv3d_v18_ncdhw_res_bf16out"]
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
9
  from ._ops import add_op_namespace_prefix, ops
10
 
11
 
12
+ @torch.library.register_fake(
13
+ add_op_namespace_prefix("bf16_causal_conv3d_ndhwc_bf16")
14
+ )
15
+ def _bf16_causal_conv3d_fake(cache_x, new_x, weight, bias, alpha, out) -> None:
16
+ n, t, h, w, ci = new_x.shape
17
+ co = weight.shape[0]
18
+ if cache_x.shape != (n, 2, h, w, ci):
19
+ raise RuntimeError("cache_x must have shape (N,2,H,W,Ci)")
20
+ if weight.shape != (co, 3, 3, 3, ci):
21
+ raise RuntimeError("weight must have shape (Co,3,3,3,Ci)")
22
+ if bias.shape != (co,) or out.shape != (n, t, h, w, co):
23
+ raise RuntimeError("bias/out shape mismatch")
24
+ return None
25
+
26
+
27
  @torch.library.register_fake(add_op_namespace_prefix("fp8_conv3d_v18_ncdhw_res_bf16out"))
28
  def _fp8_conv3d_fake(
29
  cache_x: torch.Tensor,
 
42
  raise RuntimeError("cache_x must have shape (N,2,H,W,Ci)")
43
  if weight.shape != (co, 3, 3, 3, ci):
44
  raise RuntimeError("weight must have shape (Co,3,3,3,Ci)")
45
+ if ci not in (32, 64):
46
+ raise RuntimeError(
47
+ "FP8 Conv3D is accepted only for Ci=32/64; use the "
48
+ "strong-library or NVFP4 path for larger channels"
49
+ )
50
  if residual.shape != (n, co, t_new, h, w) or out.shape != residual.shape:
51
  raise RuntimeError("residual/out must be NCDHW")
52
  if bias.shape != (co,):
 
54
  return None
55
 
56
 
57
+ @torch.library.register_fake(
58
+ add_op_namespace_prefix("fp8_causal_conv3d_ndhwc_bf16")
59
+ )
60
+ def _fp8_causal_conv3d_fake(
61
+ cache_x: torch.Tensor,
62
+ new_x: torch.Tensor,
63
+ weight: torch.Tensor,
64
+ bias: torch.Tensor,
65
+ alpha: float,
66
+ out: torch.Tensor,
67
+ ) -> None:
68
+ if cache_x.dim() != 5 or new_x.dim() != 5:
69
+ raise RuntimeError("cache_x/new_x must be NDHWC")
70
+ n, t, h, w, ci = new_x.shape
71
+ co = weight.shape[0]
72
+ if cache_x.shape != (n, 2, h, w, ci):
73
+ raise RuntimeError("cache_x must have shape (N,2,H,W,Ci)")
74
+ if weight.shape != (co, 3, 3, 3, ci):
75
+ raise RuntimeError("weight must have shape (Co,3,3,3,Ci)")
76
+ if ci not in (32, 64):
77
+ raise RuntimeError(
78
+ "FP8 Conv3D is accepted only for Ci=32/64; use the "
79
+ "strong-library or NVFP4 path for larger channels"
80
+ )
81
+ if bias.shape != (co,) or out.shape != (n, t, h, w, co):
82
+ raise RuntimeError("bias/out shape mismatch")
83
+ return None
84
+
85
+
86
+ @torch.library.register_fake(
87
+ add_op_namespace_prefix("fp8_conv2d_3x3_nhwc_bf16")
88
+ )
89
+ def _fp8_conv2d_fake(
90
+ input: torch.Tensor,
91
+ weight: torch.Tensor,
92
+ bias: torch.Tensor,
93
+ alpha: float,
94
+ out: torch.Tensor,
95
+ ) -> None:
96
+ if input.dim() != 4:
97
+ raise RuntimeError("input must have shape (N,H,W,Ci)")
98
+ n, h, w, ci = input.shape
99
+ co = weight.shape[0]
100
+ if weight.shape != (co, 3, 3, ci):
101
+ raise RuntimeError("weight must have shape (Co,3,3,Ci)")
102
+ if bias.shape != (co,) or out.shape != (n, h, w, co):
103
+ raise RuntimeError("bias/out shape mismatch")
104
+ return None
105
+
106
+
107
+ @torch.library.register_fake(
108
+ add_op_namespace_prefix("fp8_conv2d_3x3_ncdhw_bf16")
109
+ )
110
+ def _fp8_conv2d_ncdhw_fake(
111
+ input: torch.Tensor,
112
+ weight: torch.Tensor,
113
+ bias: torch.Tensor,
114
+ alpha: float,
115
+ out: torch.Tensor,
116
+ ) -> None:
117
+ if input.dim() != 5:
118
+ raise RuntimeError("input must have shape (B,T,H,W,Ci)")
119
+ b, t, h, w, ci = input.shape
120
+ co = weight.shape[0]
121
+ if weight.shape != (co, 3, 3, ci):
122
+ raise RuntimeError("weight must have shape (Co,3,3,Ci)")
123
+ if bias.shape != (co,) or out.shape != (b, co, t, h, w):
124
+ raise RuntimeError("bias/out shape mismatch")
125
+ return None
126
+
127
+
128
+ def _check_nvfp4_conv_shapes(
129
+ cache_packed: torch.Tensor,
130
+ new_packed: torch.Tensor,
131
+ weight_packed: torch.Tensor,
132
+ cache_sf: torch.Tensor,
133
+ new_sf: torch.Tensor,
134
+ weight_sf: torch.Tensor,
135
+ bias: torch.Tensor,
136
+ outer_weight: torch.Tensor | None,
137
+ ) -> tuple[int, int, int, int, int]:
138
+ if (
139
+ cache_packed.dim() != 5
140
+ or new_packed.dim() != 5
141
+ or weight_packed.dim() != 5
142
+ ):
143
+ raise RuntimeError("packed inputs must be five-dimensional")
144
+ n, t, h, w, ci_half = new_packed.shape
145
+ ci = ci_half * 2
146
+ co = weight_packed.shape[0]
147
+ if (
148
+ (ci != 64 and ci % 128 != 0)
149
+ or co % 8 != 0
150
+ or cache_packed.shape != (n, 2, h, w, ci // 2)
151
+ or weight_packed.shape != (co, 3, 3, 3, ci // 2)
152
+ or cache_sf.shape != (n, 2, h, w, ci // 16)
153
+ or new_sf.shape != (n, t, h, w, ci // 16)
154
+ or weight_sf.shape != (co, 3, 3, 3, ci // 16)
155
+ or bias.shape != (co,)
156
+ or (outer_weight is not None and outer_weight.shape != (co,))
157
+ ):
158
+ raise RuntimeError(
159
+ "NVFP4 Conv3D accepts Ci=64 or multiples of 128; "
160
+ "other shapes must use the strong-library path"
161
+ )
162
+ return n, t, h, w, co
163
+
164
+
165
+ @torch.library.register_fake(
166
+ add_op_namespace_prefix("nvfp4_causal_conv3d_ndhwc_bf16")
167
+ )
168
+ def _nvfp4_causal_conv3d_fake(
169
+ cache_packed: torch.Tensor,
170
+ new_packed: torch.Tensor,
171
+ weight_packed: torch.Tensor,
172
+ cache_sf: torch.Tensor,
173
+ new_sf: torch.Tensor,
174
+ weight_sf: torch.Tensor,
175
+ bias: torch.Tensor,
176
+ outer_weight: torch.Tensor | None,
177
+ alpha: float,
178
+ out: torch.Tensor,
179
+ ) -> None:
180
+ del alpha
181
+ n, t, h, w, co = _check_nvfp4_conv_shapes(
182
+ cache_packed,
183
+ new_packed,
184
+ weight_packed,
185
+ cache_sf,
186
+ new_sf,
187
+ weight_sf,
188
+ bias,
189
+ outer_weight,
190
+ )
191
+ if out.shape != (n, t, h, w, co):
192
+ raise RuntimeError("out must have shape (N,T,H,W,Co)")
193
+ return None
194
+
195
+
196
+ @torch.library.register_fake(
197
+ add_op_namespace_prefix(
198
+ "nvfp4_causal_conv3d_residual_ncdhw_bf16"
199
+ )
200
+ )
201
+ def _nvfp4_causal_conv3d_residual_fake(
202
+ cache_packed: torch.Tensor,
203
+ new_packed: torch.Tensor,
204
+ weight_packed: torch.Tensor,
205
+ cache_sf: torch.Tensor,
206
+ new_sf: torch.Tensor,
207
+ weight_sf: torch.Tensor,
208
+ bias: torch.Tensor,
209
+ residual: torch.Tensor,
210
+ outer_weight: torch.Tensor | None,
211
+ alpha: float,
212
+ out: torch.Tensor,
213
+ ) -> None:
214
+ del alpha
215
+ n, t, h, w, co = _check_nvfp4_conv_shapes(
216
+ cache_packed,
217
+ new_packed,
218
+ weight_packed,
219
+ cache_sf,
220
+ new_sf,
221
+ weight_sf,
222
+ bias,
223
+ outer_weight,
224
+ )
225
+ if residual.shape != (n, co, t, h, w) or out.shape != residual.shape:
226
+ raise RuntimeError("residual/out must have shape (N,Co,T,H,W)")
227
+ return None
228
+
229
+
230
  def fp8_conv3d_v18_ncdhw_res_bf16out(
231
  cache_x: torch.Tensor,
232
  new_x: torch.Tensor,
 
247
  return out
248
 
249
 
250
+ def bf16_causal_conv3d_ndhwc_bf16(
251
+ cache_x: torch.Tensor,
252
+ new_x: torch.Tensor,
253
+ weight: torch.Tensor,
254
+ bias: torch.Tensor,
255
+ alpha: float = 1.0,
256
+ *,
257
+ out: Optional[torch.Tensor] = None,
258
+ ) -> torch.Tensor:
259
+ """Experimental SM110 BF16 causal Conv3D probe.
260
+
261
+ This entry is explicit opt-in. It is not the default world-model Conv3D
262
+ backend because the current native kernel does not beat cuDNN.
263
+ """
264
+ n, t, h, w, _ = new_x.shape
265
+ if out is None:
266
+ out = torch.empty(
267
+ (n, t, h, w, weight.shape[0]),
268
+ device=new_x.device,
269
+ dtype=torch.bfloat16,
270
+ )
271
+ ops.bf16_causal_conv3d_ndhwc_bf16(
272
+ cache_x, new_x, weight, bias, float(alpha), out
273
+ )
274
+ return out
275
+
276
+
277
+ def fp8_causal_conv3d_ndhwc_bf16(
278
+ cache_x: torch.Tensor,
279
+ new_x: torch.Tensor,
280
+ weight: torch.Tensor,
281
+ bias: torch.Tensor,
282
+ alpha: float = 1.0,
283
+ *,
284
+ out: Optional[torch.Tensor] = None,
285
+ ) -> torch.Tensor:
286
+ """FP8 causal 3D convolution with virtual two-frame cache concat."""
287
+
288
+ n, t, h, w, _ = new_x.shape
289
+ if out is None:
290
+ out = torch.empty(
291
+ (n, t, h, w, weight.shape[0]),
292
+ device=new_x.device,
293
+ dtype=torch.bfloat16,
294
+ )
295
+ ops.fp8_causal_conv3d_ndhwc_bf16(
296
+ cache_x, new_x, weight, bias, float(alpha), out
297
+ )
298
+ return out
299
+
300
+
301
+ def fp8_conv2d_3x3_nhwc_bf16(
302
+ input: torch.Tensor,
303
+ weight: torch.Tensor,
304
+ bias: torch.Tensor,
305
+ alpha: float = 1.0,
306
+ *,
307
+ out: Optional[torch.Tensor] = None,
308
+ ) -> torch.Tensor:
309
+ """FP8 3x3 Conv2D with NHWC input/output and BF16 epilogue."""
310
+
311
+ if out is None:
312
+ out = torch.empty(
313
+ (*input.shape[:3], weight.shape[0]),
314
+ device=input.device,
315
+ dtype=torch.bfloat16,
316
+ )
317
+ ops.fp8_conv2d_3x3_nhwc_bf16(
318
+ input, weight, bias, float(alpha), out
319
+ )
320
+ return out
321
+
322
+
323
+ def fp8_conv2d_3x3_ncdhw_bf16(
324
+ input: torch.Tensor,
325
+ weight: torch.Tensor,
326
+ bias: torch.Tensor,
327
+ alpha: float = 1.0,
328
+ *,
329
+ out: Optional[torch.Tensor] = None,
330
+ ) -> torch.Tensor:
331
+ """FP8 3x3 Conv2D over B*T frames with direct BF16 NCDHW output."""
332
+
333
+ if out is None:
334
+ out = torch.empty(
335
+ (
336
+ input.shape[0],
337
+ weight.shape[0],
338
+ input.shape[1],
339
+ input.shape[2],
340
+ input.shape[3],
341
+ ),
342
+ device=input.device,
343
+ dtype=torch.bfloat16,
344
+ )
345
+ ops.fp8_conv2d_3x3_ncdhw_bf16(
346
+ input, weight, bias, float(alpha), out
347
+ )
348
+ return out
349
+
350
+
351
+ def nvfp4_causal_conv3d_ndhwc_bf16(
352
+ cache_packed: torch.Tensor,
353
+ new_packed: torch.Tensor,
354
+ weight_packed: torch.Tensor,
355
+ cache_sf: torch.Tensor,
356
+ new_sf: torch.Tensor,
357
+ weight_sf: torch.Tensor,
358
+ bias: torch.Tensor,
359
+ outer_weight: torch.Tensor | None = None,
360
+ alpha: float = 1.0,
361
+ *,
362
+ out: Optional[torch.Tensor] = None,
363
+ ) -> torch.Tensor:
364
+ """NVFP4 causal Conv3D with linear UE4M3 scale factors and BF16 NDHWC output."""
365
+
366
+ n, t, h, w, _ = new_packed.shape
367
+ if out is None:
368
+ out = torch.empty(
369
+ (n, t, h, w, weight_packed.shape[0]),
370
+ device=new_packed.device,
371
+ dtype=torch.bfloat16,
372
+ )
373
+ ops.nvfp4_causal_conv3d_ndhwc_bf16(
374
+ cache_packed,
375
+ new_packed,
376
+ weight_packed,
377
+ cache_sf,
378
+ new_sf,
379
+ weight_sf,
380
+ bias,
381
+ outer_weight,
382
+ float(alpha),
383
+ out,
384
+ )
385
+ return out
386
+
387
+
388
+ def nvfp4_causal_conv3d_residual_ncdhw_bf16(
389
+ cache_packed: torch.Tensor,
390
+ new_packed: torch.Tensor,
391
+ weight_packed: torch.Tensor,
392
+ cache_sf: torch.Tensor,
393
+ new_sf: torch.Tensor,
394
+ weight_sf: torch.Tensor,
395
+ bias: torch.Tensor,
396
+ residual: torch.Tensor,
397
+ outer_weight: torch.Tensor | None = None,
398
+ alpha: float = 1.0,
399
+ *,
400
+ out: Optional[torch.Tensor] = None,
401
+ ) -> torch.Tensor:
402
+ """NVFP4 causal Conv3D with fused bias, residual, and BF16 NCDHW output."""
403
+
404
+ if out is None:
405
+ out = torch.empty_like(residual)
406
+ ops.nvfp4_causal_conv3d_residual_ncdhw_bf16(
407
+ cache_packed,
408
+ new_packed,
409
+ weight_packed,
410
+ cache_sf,
411
+ new_sf,
412
+ weight_sf,
413
+ bias,
414
+ residual,
415
+ outer_weight,
416
+ float(alpha),
417
+ out,
418
+ )
419
+ return out
420
+
421
+
422
+ __all__ = [
423
+ "bf16_causal_conv3d_ndhwc_bf16",
424
+ "fp8_conv3d_v18_ncdhw_res_bf16out",
425
+ "fp8_causal_conv3d_ndhwc_bf16",
426
+ "fp8_conv2d_3x3_nhwc_bf16",
427
+ "fp8_conv2d_3x3_ncdhw_bf16",
428
+ "nvfp4_causal_conv3d_ndhwc_bf16",
429
+ "nvfp4_causal_conv3d_residual_ncdhw_bf16",
430
+ ]
build/torch212-cxx11-cu130-x86_64-linux/_ops.py CHANGED
@@ -1,9 +1,9 @@
1
  import torch
2
- from . import _world_model_conv_cuda_f14c443
3
- ops = torch.ops._world_model_conv_cuda_f14c443
4
 
5
  def add_op_namespace_prefix(op_name: str):
6
  """
7
  Prefix op by namespace.
8
  """
9
- return f"_world_model_conv_cuda_f14c443::{op_name}"
 
1
  import torch
2
+ from . import _world_model_conv_cuda_33f8494
3
+ ops = torch.ops._world_model_conv_cuda_33f8494
4
 
5
  def add_op_namespace_prefix(op_name: str):
6
  """
7
  Prefix op by namespace.
8
  """
9
+ return f"_world_model_conv_cuda_33f8494::{op_name}"
build/torch212-cxx11-cu130-x86_64-linux/{_world_model_conv_cuda_f14c443.abi3.so → _world_model_conv_cuda_33f8494.abi3.so} RENAMED
@@ -1,3 +1,3 @@
1
  version https://git-lfs.github.com/spec/v1
2
- oid sha256:1db263ef790ed8a7da8c88066e103e16e7ab7f1a8341a586fcb6cb13aa28c36b
3
- size 296000
 
1
  version https://git-lfs.github.com/spec/v1
2
+ oid sha256:d5a00c775bc1c70710c94c3de8a6ef0cb72b3616e0efaf595aabf2828469dacc
3
+ size 1300096
build/torch212-cxx11-cu130-x86_64-linux/metadata.json CHANGED
@@ -1,12 +1,13 @@
1
  {
2
  "name": "world-model-conv",
3
- "id": "_world_model_conv_cuda_f14c443",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "python-depends": [],
7
  "backend": {
8
  "type": "cuda",
9
  "archs": [
 
10
  "12.0a"
11
  ]
12
  }
 
1
  {
2
  "name": "world-model-conv",
3
+ "id": "_world_model_conv_cuda_33f8494",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "python-depends": [],
7
  "backend": {
8
  "type": "cuda",
9
  "archs": [
10
+ "11.0a",
11
  "12.0a"
12
  ]
13
  }
build/torch212-cxx11-cu132-x86_64-linux/__init__.py CHANGED
@@ -9,6 +9,21 @@ import torch
9
  from ._ops import add_op_namespace_prefix, ops
10
 
11
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
12
  @torch.library.register_fake(add_op_namespace_prefix("fp8_conv3d_v18_ncdhw_res_bf16out"))
13
  def _fp8_conv3d_fake(
14
  cache_x: torch.Tensor,
@@ -27,6 +42,11 @@ def _fp8_conv3d_fake(
27
  raise RuntimeError("cache_x must have shape (N,2,H,W,Ci)")
28
  if weight.shape != (co, 3, 3, 3, ci):
29
  raise RuntimeError("weight must have shape (Co,3,3,3,Ci)")
 
 
 
 
 
30
  if residual.shape != (n, co, t_new, h, w) or out.shape != residual.shape:
31
  raise RuntimeError("residual/out must be NCDHW")
32
  if bias.shape != (co,):
@@ -34,6 +54,179 @@ def _fp8_conv3d_fake(
34
  return None
35
 
36
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
37
  def fp8_conv3d_v18_ncdhw_res_bf16out(
38
  cache_x: torch.Tensor,
39
  new_x: torch.Tensor,
@@ -54,4 +247,184 @@ def fp8_conv3d_v18_ncdhw_res_bf16out(
54
  return out
55
 
56
 
57
- __all__ = ["fp8_conv3d_v18_ncdhw_res_bf16out"]
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
9
  from ._ops import add_op_namespace_prefix, ops
10
 
11
 
12
+ @torch.library.register_fake(
13
+ add_op_namespace_prefix("bf16_causal_conv3d_ndhwc_bf16")
14
+ )
15
+ def _bf16_causal_conv3d_fake(cache_x, new_x, weight, bias, alpha, out) -> None:
16
+ n, t, h, w, ci = new_x.shape
17
+ co = weight.shape[0]
18
+ if cache_x.shape != (n, 2, h, w, ci):
19
+ raise RuntimeError("cache_x must have shape (N,2,H,W,Ci)")
20
+ if weight.shape != (co, 3, 3, 3, ci):
21
+ raise RuntimeError("weight must have shape (Co,3,3,3,Ci)")
22
+ if bias.shape != (co,) or out.shape != (n, t, h, w, co):
23
+ raise RuntimeError("bias/out shape mismatch")
24
+ return None
25
+
26
+
27
  @torch.library.register_fake(add_op_namespace_prefix("fp8_conv3d_v18_ncdhw_res_bf16out"))
28
  def _fp8_conv3d_fake(
29
  cache_x: torch.Tensor,
 
42
  raise RuntimeError("cache_x must have shape (N,2,H,W,Ci)")
43
  if weight.shape != (co, 3, 3, 3, ci):
44
  raise RuntimeError("weight must have shape (Co,3,3,3,Ci)")
45
+ if ci not in (32, 64):
46
+ raise RuntimeError(
47
+ "FP8 Conv3D is accepted only for Ci=32/64; use the "
48
+ "strong-library or NVFP4 path for larger channels"
49
+ )
50
  if residual.shape != (n, co, t_new, h, w) or out.shape != residual.shape:
51
  raise RuntimeError("residual/out must be NCDHW")
52
  if bias.shape != (co,):
 
54
  return None
55
 
56
 
57
+ @torch.library.register_fake(
58
+ add_op_namespace_prefix("fp8_causal_conv3d_ndhwc_bf16")
59
+ )
60
+ def _fp8_causal_conv3d_fake(
61
+ cache_x: torch.Tensor,
62
+ new_x: torch.Tensor,
63
+ weight: torch.Tensor,
64
+ bias: torch.Tensor,
65
+ alpha: float,
66
+ out: torch.Tensor,
67
+ ) -> None:
68
+ if cache_x.dim() != 5 or new_x.dim() != 5:
69
+ raise RuntimeError("cache_x/new_x must be NDHWC")
70
+ n, t, h, w, ci = new_x.shape
71
+ co = weight.shape[0]
72
+ if cache_x.shape != (n, 2, h, w, ci):
73
+ raise RuntimeError("cache_x must have shape (N,2,H,W,Ci)")
74
+ if weight.shape != (co, 3, 3, 3, ci):
75
+ raise RuntimeError("weight must have shape (Co,3,3,3,Ci)")
76
+ if ci not in (32, 64):
77
+ raise RuntimeError(
78
+ "FP8 Conv3D is accepted only for Ci=32/64; use the "
79
+ "strong-library or NVFP4 path for larger channels"
80
+ )
81
+ if bias.shape != (co,) or out.shape != (n, t, h, w, co):
82
+ raise RuntimeError("bias/out shape mismatch")
83
+ return None
84
+
85
+
86
+ @torch.library.register_fake(
87
+ add_op_namespace_prefix("fp8_conv2d_3x3_nhwc_bf16")
88
+ )
89
+ def _fp8_conv2d_fake(
90
+ input: torch.Tensor,
91
+ weight: torch.Tensor,
92
+ bias: torch.Tensor,
93
+ alpha: float,
94
+ out: torch.Tensor,
95
+ ) -> None:
96
+ if input.dim() != 4:
97
+ raise RuntimeError("input must have shape (N,H,W,Ci)")
98
+ n, h, w, ci = input.shape
99
+ co = weight.shape[0]
100
+ if weight.shape != (co, 3, 3, ci):
101
+ raise RuntimeError("weight must have shape (Co,3,3,Ci)")
102
+ if bias.shape != (co,) or out.shape != (n, h, w, co):
103
+ raise RuntimeError("bias/out shape mismatch")
104
+ return None
105
+
106
+
107
+ @torch.library.register_fake(
108
+ add_op_namespace_prefix("fp8_conv2d_3x3_ncdhw_bf16")
109
+ )
110
+ def _fp8_conv2d_ncdhw_fake(
111
+ input: torch.Tensor,
112
+ weight: torch.Tensor,
113
+ bias: torch.Tensor,
114
+ alpha: float,
115
+ out: torch.Tensor,
116
+ ) -> None:
117
+ if input.dim() != 5:
118
+ raise RuntimeError("input must have shape (B,T,H,W,Ci)")
119
+ b, t, h, w, ci = input.shape
120
+ co = weight.shape[0]
121
+ if weight.shape != (co, 3, 3, ci):
122
+ raise RuntimeError("weight must have shape (Co,3,3,Ci)")
123
+ if bias.shape != (co,) or out.shape != (b, co, t, h, w):
124
+ raise RuntimeError("bias/out shape mismatch")
125
+ return None
126
+
127
+
128
+ def _check_nvfp4_conv_shapes(
129
+ cache_packed: torch.Tensor,
130
+ new_packed: torch.Tensor,
131
+ weight_packed: torch.Tensor,
132
+ cache_sf: torch.Tensor,
133
+ new_sf: torch.Tensor,
134
+ weight_sf: torch.Tensor,
135
+ bias: torch.Tensor,
136
+ outer_weight: torch.Tensor | None,
137
+ ) -> tuple[int, int, int, int, int]:
138
+ if (
139
+ cache_packed.dim() != 5
140
+ or new_packed.dim() != 5
141
+ or weight_packed.dim() != 5
142
+ ):
143
+ raise RuntimeError("packed inputs must be five-dimensional")
144
+ n, t, h, w, ci_half = new_packed.shape
145
+ ci = ci_half * 2
146
+ co = weight_packed.shape[0]
147
+ if (
148
+ (ci != 64 and ci % 128 != 0)
149
+ or co % 8 != 0
150
+ or cache_packed.shape != (n, 2, h, w, ci // 2)
151
+ or weight_packed.shape != (co, 3, 3, 3, ci // 2)
152
+ or cache_sf.shape != (n, 2, h, w, ci // 16)
153
+ or new_sf.shape != (n, t, h, w, ci // 16)
154
+ or weight_sf.shape != (co, 3, 3, 3, ci // 16)
155
+ or bias.shape != (co,)
156
+ or (outer_weight is not None and outer_weight.shape != (co,))
157
+ ):
158
+ raise RuntimeError(
159
+ "NVFP4 Conv3D accepts Ci=64 or multiples of 128; "
160
+ "other shapes must use the strong-library path"
161
+ )
162
+ return n, t, h, w, co
163
+
164
+
165
+ @torch.library.register_fake(
166
+ add_op_namespace_prefix("nvfp4_causal_conv3d_ndhwc_bf16")
167
+ )
168
+ def _nvfp4_causal_conv3d_fake(
169
+ cache_packed: torch.Tensor,
170
+ new_packed: torch.Tensor,
171
+ weight_packed: torch.Tensor,
172
+ cache_sf: torch.Tensor,
173
+ new_sf: torch.Tensor,
174
+ weight_sf: torch.Tensor,
175
+ bias: torch.Tensor,
176
+ outer_weight: torch.Tensor | None,
177
+ alpha: float,
178
+ out: torch.Tensor,
179
+ ) -> None:
180
+ del alpha
181
+ n, t, h, w, co = _check_nvfp4_conv_shapes(
182
+ cache_packed,
183
+ new_packed,
184
+ weight_packed,
185
+ cache_sf,
186
+ new_sf,
187
+ weight_sf,
188
+ bias,
189
+ outer_weight,
190
+ )
191
+ if out.shape != (n, t, h, w, co):
192
+ raise RuntimeError("out must have shape (N,T,H,W,Co)")
193
+ return None
194
+
195
+
196
+ @torch.library.register_fake(
197
+ add_op_namespace_prefix(
198
+ "nvfp4_causal_conv3d_residual_ncdhw_bf16"
199
+ )
200
+ )
201
+ def _nvfp4_causal_conv3d_residual_fake(
202
+ cache_packed: torch.Tensor,
203
+ new_packed: torch.Tensor,
204
+ weight_packed: torch.Tensor,
205
+ cache_sf: torch.Tensor,
206
+ new_sf: torch.Tensor,
207
+ weight_sf: torch.Tensor,
208
+ bias: torch.Tensor,
209
+ residual: torch.Tensor,
210
+ outer_weight: torch.Tensor | None,
211
+ alpha: float,
212
+ out: torch.Tensor,
213
+ ) -> None:
214
+ del alpha
215
+ n, t, h, w, co = _check_nvfp4_conv_shapes(
216
+ cache_packed,
217
+ new_packed,
218
+ weight_packed,
219
+ cache_sf,
220
+ new_sf,
221
+ weight_sf,
222
+ bias,
223
+ outer_weight,
224
+ )
225
+ if residual.shape != (n, co, t, h, w) or out.shape != residual.shape:
226
+ raise RuntimeError("residual/out must have shape (N,Co,T,H,W)")
227
+ return None
228
+
229
+
230
  def fp8_conv3d_v18_ncdhw_res_bf16out(
231
  cache_x: torch.Tensor,
232
  new_x: torch.Tensor,
 
247
  return out
248
 
249
 
250
+ def bf16_causal_conv3d_ndhwc_bf16(
251
+ cache_x: torch.Tensor,
252
+ new_x: torch.Tensor,
253
+ weight: torch.Tensor,
254
+ bias: torch.Tensor,
255
+ alpha: float = 1.0,
256
+ *,
257
+ out: Optional[torch.Tensor] = None,
258
+ ) -> torch.Tensor:
259
+ """Experimental SM110 BF16 causal Conv3D probe.
260
+
261
+ This entry is explicit opt-in. It is not the default world-model Conv3D
262
+ backend because the current native kernel does not beat cuDNN.
263
+ """
264
+ n, t, h, w, _ = new_x.shape
265
+ if out is None:
266
+ out = torch.empty(
267
+ (n, t, h, w, weight.shape[0]),
268
+ device=new_x.device,
269
+ dtype=torch.bfloat16,
270
+ )
271
+ ops.bf16_causal_conv3d_ndhwc_bf16(
272
+ cache_x, new_x, weight, bias, float(alpha), out
273
+ )
274
+ return out
275
+
276
+
277
+ def fp8_causal_conv3d_ndhwc_bf16(
278
+ cache_x: torch.Tensor,
279
+ new_x: torch.Tensor,
280
+ weight: torch.Tensor,
281
+ bias: torch.Tensor,
282
+ alpha: float = 1.0,
283
+ *,
284
+ out: Optional[torch.Tensor] = None,
285
+ ) -> torch.Tensor:
286
+ """FP8 causal 3D convolution with virtual two-frame cache concat."""
287
+
288
+ n, t, h, w, _ = new_x.shape
289
+ if out is None:
290
+ out = torch.empty(
291
+ (n, t, h, w, weight.shape[0]),
292
+ device=new_x.device,
293
+ dtype=torch.bfloat16,
294
+ )
295
+ ops.fp8_causal_conv3d_ndhwc_bf16(
296
+ cache_x, new_x, weight, bias, float(alpha), out
297
+ )
298
+ return out
299
+
300
+
301
+ def fp8_conv2d_3x3_nhwc_bf16(
302
+ input: torch.Tensor,
303
+ weight: torch.Tensor,
304
+ bias: torch.Tensor,
305
+ alpha: float = 1.0,
306
+ *,
307
+ out: Optional[torch.Tensor] = None,
308
+ ) -> torch.Tensor:
309
+ """FP8 3x3 Conv2D with NHWC input/output and BF16 epilogue."""
310
+
311
+ if out is None:
312
+ out = torch.empty(
313
+ (*input.shape[:3], weight.shape[0]),
314
+ device=input.device,
315
+ dtype=torch.bfloat16,
316
+ )
317
+ ops.fp8_conv2d_3x3_nhwc_bf16(
318
+ input, weight, bias, float(alpha), out
319
+ )
320
+ return out
321
+
322
+
323
+ def fp8_conv2d_3x3_ncdhw_bf16(
324
+ input: torch.Tensor,
325
+ weight: torch.Tensor,
326
+ bias: torch.Tensor,
327
+ alpha: float = 1.0,
328
+ *,
329
+ out: Optional[torch.Tensor] = None,
330
+ ) -> torch.Tensor:
331
+ """FP8 3x3 Conv2D over B*T frames with direct BF16 NCDHW output."""
332
+
333
+ if out is None:
334
+ out = torch.empty(
335
+ (
336
+ input.shape[0],
337
+ weight.shape[0],
338
+ input.shape[1],
339
+ input.shape[2],
340
+ input.shape[3],
341
+ ),
342
+ device=input.device,
343
+ dtype=torch.bfloat16,
344
+ )
345
+ ops.fp8_conv2d_3x3_ncdhw_bf16(
346
+ input, weight, bias, float(alpha), out
347
+ )
348
+ return out
349
+
350
+
351
+ def nvfp4_causal_conv3d_ndhwc_bf16(
352
+ cache_packed: torch.Tensor,
353
+ new_packed: torch.Tensor,
354
+ weight_packed: torch.Tensor,
355
+ cache_sf: torch.Tensor,
356
+ new_sf: torch.Tensor,
357
+ weight_sf: torch.Tensor,
358
+ bias: torch.Tensor,
359
+ outer_weight: torch.Tensor | None = None,
360
+ alpha: float = 1.0,
361
+ *,
362
+ out: Optional[torch.Tensor] = None,
363
+ ) -> torch.Tensor:
364
+ """NVFP4 causal Conv3D with linear UE4M3 scale factors and BF16 NDHWC output."""
365
+
366
+ n, t, h, w, _ = new_packed.shape
367
+ if out is None:
368
+ out = torch.empty(
369
+ (n, t, h, w, weight_packed.shape[0]),
370
+ device=new_packed.device,
371
+ dtype=torch.bfloat16,
372
+ )
373
+ ops.nvfp4_causal_conv3d_ndhwc_bf16(
374
+ cache_packed,
375
+ new_packed,
376
+ weight_packed,
377
+ cache_sf,
378
+ new_sf,
379
+ weight_sf,
380
+ bias,
381
+ outer_weight,
382
+ float(alpha),
383
+ out,
384
+ )
385
+ return out
386
+
387
+
388
+ def nvfp4_causal_conv3d_residual_ncdhw_bf16(
389
+ cache_packed: torch.Tensor,
390
+ new_packed: torch.Tensor,
391
+ weight_packed: torch.Tensor,
392
+ cache_sf: torch.Tensor,
393
+ new_sf: torch.Tensor,
394
+ weight_sf: torch.Tensor,
395
+ bias: torch.Tensor,
396
+ residual: torch.Tensor,
397
+ outer_weight: torch.Tensor | None = None,
398
+ alpha: float = 1.0,
399
+ *,
400
+ out: Optional[torch.Tensor] = None,
401
+ ) -> torch.Tensor:
402
+ """NVFP4 causal Conv3D with fused bias, residual, and BF16 NCDHW output."""
403
+
404
+ if out is None:
405
+ out = torch.empty_like(residual)
406
+ ops.nvfp4_causal_conv3d_residual_ncdhw_bf16(
407
+ cache_packed,
408
+ new_packed,
409
+ weight_packed,
410
+ cache_sf,
411
+ new_sf,
412
+ weight_sf,
413
+ bias,
414
+ residual,
415
+ outer_weight,
416
+ float(alpha),
417
+ out,
418
+ )
419
+ return out
420
+
421
+
422
+ __all__ = [
423
+ "bf16_causal_conv3d_ndhwc_bf16",
424
+ "fp8_conv3d_v18_ncdhw_res_bf16out",
425
+ "fp8_causal_conv3d_ndhwc_bf16",
426
+ "fp8_conv2d_3x3_nhwc_bf16",
427
+ "fp8_conv2d_3x3_ncdhw_bf16",
428
+ "nvfp4_causal_conv3d_ndhwc_bf16",
429
+ "nvfp4_causal_conv3d_residual_ncdhw_bf16",
430
+ ]
build/torch212-cxx11-cu132-x86_64-linux/_ops.py CHANGED
@@ -1,9 +1,9 @@
1
  import torch
2
- from . import _world_model_conv_cuda_f14c443
3
- ops = torch.ops._world_model_conv_cuda_f14c443
4
 
5
  def add_op_namespace_prefix(op_name: str):
6
  """
7
  Prefix op by namespace.
8
  """
9
- return f"_world_model_conv_cuda_f14c443::{op_name}"
 
1
  import torch
2
+ from . import _world_model_conv_cuda_33f8494
3
+ ops = torch.ops._world_model_conv_cuda_33f8494
4
 
5
  def add_op_namespace_prefix(op_name: str):
6
  """
7
  Prefix op by namespace.
8
  """
9
+ return f"_world_model_conv_cuda_33f8494::{op_name}"
build/torch212-cxx11-cu132-x86_64-linux/_world_model_conv_cuda_33f8494.abi3.so ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:3e28e6b612c7ea5e752218cbf591c8470fa524dc6b85ef21cdba6a41c2a29f05
3
+ size 1336960
build/torch212-cxx11-cu132-x86_64-linux/metadata.json CHANGED
@@ -1,12 +1,13 @@
1
  {
2
  "name": "world-model-conv",
3
- "id": "_world_model_conv_cuda_f14c443",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "python-depends": [],
7
  "backend": {
8
  "type": "cuda",
9
  "archs": [
 
10
  "12.0a"
11
  ]
12
  }
 
1
  {
2
  "name": "world-model-conv",
3
+ "id": "_world_model_conv_cuda_33f8494",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "python-depends": [],
7
  "backend": {
8
  "type": "cuda",
9
  "archs": [
10
+ "11.0a",
11
  "12.0a"
12
  ]
13
  }