flash-rt 0.2.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (455) hide show
  1. flash_rt/__init__.py +107 -0
  2. flash_rt/_extensions.py +119 -0
  3. flash_rt/amd/__init__.py +9 -0
  4. flash_rt/amd/core/__init__.py +0 -0
  5. flash_rt/amd/core/hip_buffer.py +176 -0
  6. flash_rt/amd/core/hip_graph.py +102 -0
  7. flash_rt/amd/frontends/__init__.py +1 -0
  8. flash_rt/amd/frontends/torch/__init__.py +1 -0
  9. flash_rt/amd/frontends/torch/groot_n17.py +1013 -0
  10. flash_rt/amd/frontends/torch/pi05.py +1431 -0
  11. flash_rt/amd/hardware/__init__.py +1 -0
  12. flash_rt/amd/hardware/cdna4/__init__.py +1 -0
  13. flash_rt/amd/hardware/cdna4/attn_backend.py +338 -0
  14. flash_rt/amd/hardware/cdna4/attn_backend_aiter.py +410 -0
  15. flash_rt/amd/hardware/cdna4/attn_backend_groot_n17.py +297 -0
  16. flash_rt/amd/models/__init__.py +1 -0
  17. flash_rt/amd/models/groot_n17/__init__.py +1 -0
  18. flash_rt/amd/models/groot_n17/pipeline.py +1058 -0
  19. flash_rt/amd/models/pi05/__init__.py +1 -0
  20. flash_rt/amd/models/pi05/pipeline.py +1860 -0
  21. flash_rt/api.py +1144 -0
  22. flash_rt/catalog/__init__.py +39 -0
  23. flash_rt/catalog/binding.py +412 -0
  24. flash_rt/catalog/bindings/cosmos3_video_pipeline.yaml +94 -0
  25. flash_rt/catalog/bindings/groot_n16_dit.yaml +23 -0
  26. flash_rt/catalog/bindings/groot_n16_llm.yaml +21 -0
  27. flash_rt/catalog/bindings/groot_n16_pipeline.yaml +107 -0
  28. flash_rt/catalog/bindings/groot_n16_tick.yaml +26 -0
  29. flash_rt/catalog/bindings/groot_n16_vision.yaml +23 -0
  30. flash_rt/catalog/bindings/groot_n17_pipeline.yaml +117 -0
  31. flash_rt/catalog/bindings/lingbot_vla_pipeline.yaml +106 -0
  32. flash_rt/catalog/bindings/motus_tick.yaml +120 -0
  33. flash_rt/catalog/bindings/nexn2_pipeline.yaml +112 -0
  34. flash_rt/catalog/bindings/pi05.yaml +29 -0
  35. flash_rt/catalog/bindings/pi05_prefix.yaml +21 -0
  36. flash_rt/catalog/bindings/pi05_tick.yaml +94 -0
  37. flash_rt/catalog/bindings/pi05_vision.yaml +23 -0
  38. flash_rt/catalog/bindings/qwen25_15b.yaml +22 -0
  39. flash_rt/catalog/bindings/qwen36_27b_pipeline.yaml +130 -0
  40. flash_rt/catalog/bindings/qwen3_8b.yaml +22 -0
  41. flash_rt/catalog/bindings/qwen3_8b_pipeline.yaml +98 -0
  42. flash_rt/catalog/bindings/qwen3_vl_8b_pipeline.yaml +131 -0
  43. flash_rt/catalog/bindings/qwen3_vl_8b_text.yaml +22 -0
  44. flash_rt/catalog/bindings/qwen3_vl_8b_vision.yaml +23 -0
  45. flash_rt/catalog/bindings/smolvla_base.yaml +21 -0
  46. flash_rt/catalog/bindings/smolvla_expert.yaml +21 -0
  47. flash_rt/catalog/bindings/smolvla_pipeline.yaml +93 -0
  48. flash_rt/catalog/bindings/smolvla_tick.yaml +27 -0
  49. flash_rt/catalog/bindings/smolvla_vision.yaml +23 -0
  50. flash_rt/catalog/bindings/wan22_video_pipeline.yaml +119 -0
  51. flash_rt/catalog/registry.py +108 -0
  52. flash_rt/catalog/structures/__init__.py +0 -0
  53. flash_rt/catalog/structures/adaln_producer/__init__.py +0 -0
  54. flash_rt/catalog/structures/adaln_producer/reference.py +40 -0
  55. flash_rt/catalog/structures/adaln_producer/structure.yaml +106 -0
  56. flash_rt/catalog/structures/attention_core/__init__.py +0 -0
  57. flash_rt/catalog/structures/attention_core/reference.py +36 -0
  58. flash_rt/catalog/structures/attention_core/structure.yaml +72 -0
  59. flash_rt/catalog/structures/autoregressive_decode_pipeline/structure.yaml +82 -0
  60. flash_rt/catalog/structures/cadence_static/__init__.py +0 -0
  61. flash_rt/catalog/structures/cadence_static/reference.py +24 -0
  62. flash_rt/catalog/structures/cadence_static/structure.yaml +56 -0
  63. flash_rt/catalog/structures/decoder_block/__init__.py +0 -0
  64. flash_rt/catalog/structures/decoder_block/reference.py +38 -0
  65. flash_rt/catalog/structures/decoder_block/structure.yaml +86 -0
  66. flash_rt/catalog/structures/decoder_ffn/__init__.py +0 -0
  67. flash_rt/catalog/structures/decoder_ffn/reference.py +64 -0
  68. flash_rt/catalog/structures/decoder_ffn/structure.yaml +44 -0
  69. flash_rt/catalog/structures/gated_delta_core/reference.py +54 -0
  70. flash_rt/catalog/structures/gated_delta_core/structure.yaml +60 -0
  71. flash_rt/catalog/structures/linear_proj/__init__.py +3 -0
  72. flash_rt/catalog/structures/linear_proj/reference.py +35 -0
  73. flash_rt/catalog/structures/linear_proj/structure.yaml +74 -0
  74. flash_rt/catalog/structures/modnorm_qkv_chain/__init__.py +1 -0
  75. flash_rt/catalog/structures/modnorm_qkv_chain/reference.py +39 -0
  76. flash_rt/catalog/structures/modnorm_qkv_chain/structure.yaml +59 -0
  77. flash_rt/catalog/structures/norm_fused/__init__.py +0 -0
  78. flash_rt/catalog/structures/norm_fused/reference.py +26 -0
  79. flash_rt/catalog/structures/norm_fused/structure.yaml +50 -0
  80. flash_rt/catalog/structures/patch_projection/reference.py +15 -0
  81. flash_rt/catalog/structures/patch_projection/structure.yaml +53 -0
  82. flash_rt/catalog/structures/qk_norm_rope/__init__.py +3 -0
  83. flash_rt/catalog/structures/qk_norm_rope/reference.py +114 -0
  84. flash_rt/catalog/structures/qk_norm_rope/structure.yaml +83 -0
  85. flash_rt/catalog/structures/qkv_pack/__init__.py +0 -0
  86. flash_rt/catalog/structures/qkv_pack/reference.py +33 -0
  87. flash_rt/catalog/structures/qkv_pack/structure.yaml +65 -0
  88. flash_rt/catalog/structures/qkv_rope/__init__.py +1 -0
  89. flash_rt/catalog/structures/qkv_rope/reference.py +39 -0
  90. flash_rt/catalog/structures/qkv_rope/structure.yaml +55 -0
  91. flash_rt/catalog/structures/video_generation_pipeline/__init__.py +2 -0
  92. flash_rt/catalog/structures/video_generation_pipeline/structure.yaml +81 -0
  93. flash_rt/catalog/structures/vision_ffn/__init__.py +0 -0
  94. flash_rt/catalog/structures/vision_ffn/reference.py +35 -0
  95. flash_rt/catalog/structures/vision_ffn/structure.yaml +42 -0
  96. flash_rt/catalog/structures/vla_tick_pipeline/__init__.py +7 -0
  97. flash_rt/catalog/structures/vla_tick_pipeline/structure.yaml +82 -0
  98. flash_rt/configs/__init__.py +0 -0
  99. flash_rt/configs/cosmos3_edge.yaml +21 -0
  100. flash_rt/configs/cosmos3_video.yaml +24 -0
  101. flash_rt/configs/groot.yaml +73 -0
  102. flash_rt/configs/groot_n17.yaml +53 -0
  103. flash_rt/configs/hyvla.yaml +65 -0
  104. flash_rt/configs/ltx25.yaml +41 -0
  105. flash_rt/configs/motus.yaml +85 -0
  106. flash_rt/configs/nexn2.yaml +79 -0
  107. flash_rt/configs/pi0.yaml +38 -0
  108. flash_rt/configs/pi05.yaml +38 -0
  109. flash_rt/configs/qwen36.yaml +68 -0
  110. flash_rt/configs/wan22_ti2v_5b.yaml +24 -0
  111. flash_rt/core/__init__.py +0 -0
  112. flash_rt/core/calibration.py +301 -0
  113. flash_rt/core/calibration_api.py +70 -0
  114. flash_rt/core/config.py +96 -0
  115. flash_rt/core/context.py +47 -0
  116. flash_rt/core/cuda_buffer.py +189 -0
  117. flash_rt/core/cuda_graph.py +81 -0
  118. flash_rt/core/parity.py +37 -0
  119. flash_rt/core/precision_spec.py +164 -0
  120. flash_rt/core/quant/__init__.py +0 -0
  121. flash_rt/core/quant/calibrator.py +170 -0
  122. flash_rt/core/quantization.py +73 -0
  123. flash_rt/core/rl/__init__.py +75 -0
  124. flash_rt/core/rl/acp_tags.py +51 -0
  125. flash_rt/core/rl/advantage.py +163 -0
  126. flash_rt/core/rl/cfg_sampler.py +72 -0
  127. flash_rt/core/rl/reward.py +233 -0
  128. flash_rt/core/rl/value_function.py +198 -0
  129. flash_rt/core/thor_frontend_utils.py +152 -0
  130. flash_rt/core/utils/__init__.py +0 -0
  131. flash_rt/core/utils/actions.py +19 -0
  132. flash_rt/core/utils/hardware.py +50 -0
  133. flash_rt/core/utils/norm_stats.py +359 -0
  134. flash_rt/core/utils/pi05_prompt.py +35 -0
  135. flash_rt/core/weights/__init__.py +0 -0
  136. flash_rt/core/weights/loader.py +135 -0
  137. flash_rt/core/weights/transformer.py +691 -0
  138. flash_rt/core/weights/weight_cache.py +147 -0
  139. flash_rt/datasets/__init__.py +11 -0
  140. flash_rt/datasets/libero.py +306 -0
  141. flash_rt/executors/__init__.py +6 -0
  142. flash_rt/executors/fp4_utils.py +241 -0
  143. flash_rt/executors/fp4_utils_cb.py +207 -0
  144. flash_rt/executors/jax_weights.py +270 -0
  145. flash_rt/executors/torch_weights.py +500 -0
  146. flash_rt/executors/weight_loader.py +331 -0
  147. flash_rt/frontends/__init__.py +8 -0
  148. flash_rt/frontends/_fp8_layout.py +32 -0
  149. flash_rt/frontends/jax/__init__.py +1 -0
  150. flash_rt/frontends/jax/_pi05_thor_spec.py +52 -0
  151. flash_rt/frontends/jax/_pi0_thor_spec.py +32 -0
  152. flash_rt/frontends/jax/_thor_spec_common.py +114 -0
  153. flash_rt/frontends/jax/pi05_rtx.py +576 -0
  154. flash_rt/frontends/jax/pi05_thor.py +2768 -0
  155. flash_rt/frontends/jax/pi05_thor_fp4.py +879 -0
  156. flash_rt/frontends/jax/pi0_rtx.py +483 -0
  157. flash_rt/frontends/jax/pi0_thor.py +1425 -0
  158. flash_rt/frontends/jax/pi0fast.py +1337 -0
  159. flash_rt/frontends/jetson_pi/__init__.py +12 -0
  160. flash_rt/frontends/jetson_pi/llm.py +261 -0
  161. flash_rt/frontends/jetson_pi/mllm.py +289 -0
  162. flash_rt/frontends/jetson_pi/pi0.py +420 -0
  163. flash_rt/frontends/torch/__init__.py +1 -0
  164. flash_rt/frontends/torch/_chameleon_quant.py +251 -0
  165. flash_rt/frontends/torch/_chameleon_rtx_sm87_spec.py +85 -0
  166. flash_rt/frontends/torch/_chameleon_thor_spec.py +92 -0
  167. flash_rt/frontends/torch/_cosmos3_edge_thor_spec.py +177 -0
  168. flash_rt/frontends/torch/_groot_n17_rtx_spec.py +13 -0
  169. flash_rt/frontends/torch/_groot_n17_thor_spec.py +406 -0
  170. flash_rt/frontends/torch/_groot_thor_spec.py +105 -0
  171. flash_rt/frontends/torch/_higgs_audio_v3_bf16.py +374 -0
  172. flash_rt/frontends/torch/_higgs_audio_v3_fp8.py +464 -0
  173. flash_rt/frontends/torch/_hyvla_thor_spec.py +196 -0
  174. flash_rt/frontends/torch/_lingbot_thor_spec.py +321 -0
  175. flash_rt/frontends/torch/_motus_rtx_spec.py +47 -0
  176. flash_rt/frontends/torch/_nexn2_rtx_decode.py +1621 -0
  177. flash_rt/frontends/torch/_nexn2_rtx_forward.py +1815 -0
  178. flash_rt/frontends/torch/_nexn2_rtx_nvfp4_weights.py +416 -0
  179. flash_rt/frontends/torch/_pi05_thor_spec.py +100 -0
  180. flash_rt/frontends/torch/_pi0_thor_spec.py +81 -0
  181. flash_rt/frontends/torch/_qwen36_rtx_dflash_forward.py +940 -0
  182. flash_rt/frontends/torch/_qwen36_rtx_dflash_weights.py +396 -0
  183. flash_rt/frontends/torch/_qwen36_rtx_nvfp4_weights.py +741 -0
  184. flash_rt/frontends/torch/_qwen36_rtx_turboquant.py +872 -0
  185. flash_rt/frontends/torch/_qwen36_rtx_weights.py +411 -0
  186. flash_rt/frontends/torch/_qwen3_rtx_nvfp4_weights.py +575 -0
  187. flash_rt/frontends/torch/_qwen3_vl_bf16_weights.py +191 -0
  188. flash_rt/frontends/torch/_qwen3_vl_fp8_weights.py +245 -0
  189. flash_rt/frontends/torch/_qwen3_vl_geometry.py +337 -0
  190. flash_rt/frontends/torch/_qwen3_vl_vision_rtx.py +625 -0
  191. flash_rt/frontends/torch/_template/attention.py +124 -0
  192. flash_rt/frontends/torch/_template/frontend.py +330 -0
  193. flash_rt/frontends/torch/_template/pipeline.py +263 -0
  194. flash_rt/frontends/torch/_template/weights_spec.py +215 -0
  195. flash_rt/frontends/torch/_thor_spec_common.py +147 -0
  196. flash_rt/frontends/torch/chameleon_rtx_sm87.py +721 -0
  197. flash_rt/frontends/torch/chameleon_thor.py +911 -0
  198. flash_rt/frontends/torch/cosmos3_edge_thor.py +571 -0
  199. flash_rt/frontends/torch/cosmos3_video_rtx.py +130 -0
  200. flash_rt/frontends/torch/groot_n17_rtx.py +152 -0
  201. flash_rt/frontends/torch/groot_n17_rtx_fp16.py +655 -0
  202. flash_rt/frontends/torch/groot_n17_rtx_fp8.py +582 -0
  203. flash_rt/frontends/torch/groot_n17_rtx_sm89.py +609 -0
  204. flash_rt/frontends/torch/groot_n17_rtx_sm89_fp16.py +652 -0
  205. flash_rt/frontends/torch/groot_n17_thor.py +1965 -0
  206. flash_rt/frontends/torch/groot_n17_thor_fp16.py +49 -0
  207. flash_rt/frontends/torch/groot_n17_thor_fp4.py +165 -0
  208. flash_rt/frontends/torch/groot_n17_thor_fp8.py +780 -0
  209. flash_rt/frontends/torch/groot_rtx.py +1876 -0
  210. flash_rt/frontends/torch/groot_rtx_fp16.py +1162 -0
  211. flash_rt/frontends/torch/groot_thor.py +3623 -0
  212. flash_rt/frontends/torch/groot_thor_fp16.py +28 -0
  213. flash_rt/frontends/torch/higgs_audio_v3_rtx.py +601 -0
  214. flash_rt/frontends/torch/hyvla_orin.py +306 -0
  215. flash_rt/frontends/torch/hyvla_thor.py +683 -0
  216. flash_rt/frontends/torch/lingbot_thor.py +116 -0
  217. flash_rt/frontends/torch/ltx25_rtx.py +378 -0
  218. flash_rt/frontends/torch/motus_rtx.py +1562 -0
  219. flash_rt/frontends/torch/nexn2_rtx.py +310 -0
  220. flash_rt/frontends/torch/pi05_rtx.py +1949 -0
  221. flash_rt/frontends/torch/pi05_rtx_fp16.py +1806 -0
  222. flash_rt/frontends/torch/pi05_thor.py +3083 -0
  223. flash_rt/frontends/torch/pi05_thor_fp4.py +1500 -0
  224. flash_rt/frontends/torch/pi0_rtx.py +957 -0
  225. flash_rt/frontends/torch/pi0_thor.py +1405 -0
  226. flash_rt/frontends/torch/pi0fast.py +1402 -0
  227. flash_rt/frontends/torch/qwen36_moe.py +426 -0
  228. flash_rt/frontends/torch/qwen36_moe_rtx.py +20 -0
  229. flash_rt/frontends/torch/qwen36_rtx.py +12388 -0
  230. flash_rt/frontends/torch/qwen36_spark.py +200 -0
  231. flash_rt/frontends/torch/qwen36_thor.py +1320 -0
  232. flash_rt/frontends/torch/qwen3_rtx.py +2166 -0
  233. flash_rt/frontends/torch/qwen3_vl_fp8_sm89.py +912 -0
  234. flash_rt/frontends/torch/qwen3_vl_fp8_sm89_multimodal.py +456 -0
  235. flash_rt/frontends/torch/qwen3_vl_rtx.py +642 -0
  236. flash_rt/frontends/torch/qwen3_vl_rtx_bf16.py +1031 -0
  237. flash_rt/frontends/torch/qwen3_vl_thor.py +856 -0
  238. flash_rt/frontends/torch/wan22_rtx.py +477 -0
  239. flash_rt/hardware/__init__.py +311 -0
  240. flash_rt/hardware/backend.py +407 -0
  241. flash_rt/hardware/blackwell/__init__.py +12 -0
  242. flash_rt/hardware/rtx/__init__.py +37 -0
  243. flash_rt/hardware/rtx/attn_backend.py +856 -0
  244. flash_rt/hardware/rtx/attn_backend_batched_pi05.py +303 -0
  245. flash_rt/hardware/rtx/attn_backend_chameleon.py +237 -0
  246. flash_rt/hardware/rtx/attn_backend_groot.py +447 -0
  247. flash_rt/hardware/rtx/attn_backend_groot_n17.py +251 -0
  248. flash_rt/hardware/rtx/attn_backend_groot_n17_backbone.py +191 -0
  249. flash_rt/hardware/rtx/attn_backend_motus.py +128 -0
  250. flash_rt/hardware/rtx/attn_backend_nexn2.py +409 -0
  251. flash_rt/hardware/rtx/attn_backend_qwen3.py +520 -0
  252. flash_rt/hardware/rtx/attn_backend_qwen36.py +272 -0
  253. flash_rt/hardware/thor/__init__.py +9 -0
  254. flash_rt/hardware/thor/attn_backend.py +559 -0
  255. flash_rt/hardware/thor/attn_backend_chameleon.py +362 -0
  256. flash_rt/hardware/thor/attn_backend_groot.py +328 -0
  257. flash_rt/hardware/thor/attn_backend_groot_n17.py +423 -0
  258. flash_rt/hardware/thor/attn_backend_qwen3.py +229 -0
  259. flash_rt/hardware/thor/attn_backend_qwen36.py +530 -0
  260. flash_rt/hardware/thor/fa4_backend.py +117 -0
  261. flash_rt/hardware/thor/shared_primitives.py +727 -0
  262. flash_rt/hardware/thor/shared_primitives_batched.py +178 -0
  263. flash_rt/hardware/thor/shared_primitives_fp4.py +512 -0
  264. flash_rt/hardware/thor/vqgan_trt_backend.py +187 -0
  265. flash_rt/models/__init__.py +11 -0
  266. flash_rt/models/chameleon/__init__.py +18 -0
  267. flash_rt/models/chameleon/pipeline_rtx.py +305 -0
  268. flash_rt/models/chameleon/pipeline_thor.py +1126 -0
  269. flash_rt/models/chameleon/vqvae_hf.py +124 -0
  270. flash_rt/models/cosmos3_edge/__init__.py +37 -0
  271. flash_rt/models/cosmos3_edge/action_only_official.py +3447 -0
  272. flash_rt/models/cosmos3_edge/boundary_dump.py +151 -0
  273. flash_rt/models/cosmos3_edge/denoise_ref.py +325 -0
  274. flash_rt/models/cosmos3_edge/dump_replay.py +195 -0
  275. flash_rt/models/cosmos3_edge/layer_ref.py +1040 -0
  276. flash_rt/models/cosmos3_edge/pipeline_thor.py +549 -0
  277. flash_rt/models/cosmos3_edge/static_engine.py +346 -0
  278. flash_rt/models/cosmos3_edge/static_unipc.py +234 -0
  279. flash_rt/models/cosmos3_edge/vae_native.py +304 -0
  280. flash_rt/models/cosmos3_edge/weights.py +91 -0
  281. flash_rt/models/cosmos3_reasoner/__init__.py +1 -0
  282. flash_rt/models/cosmos3_reasoner/pipeline_thor.py +691 -0
  283. flash_rt/models/cosmos3_video/__init__.py +6 -0
  284. flash_rt/models/cosmos3_video/fm_solvers_unipc.py +808 -0
  285. flash_rt/models/cosmos3_video/kernels/__init__.py +22 -0
  286. flash_rt/models/cosmos3_video/kernels/csrc/bindings.cpp +17 -0
  287. flash_rt/models/cosmos3_video/kernels/csrc/fused_qk_norm_rope.cu +61 -0
  288. flash_rt/models/cosmos3_video/kernels/setup.py +33 -0
  289. flash_rt/models/cosmos3_video/pipeline_rtx.py +234 -0
  290. flash_rt/models/groot/__init__.py +32 -0
  291. flash_rt/models/groot/embodiments.py +69 -0
  292. flash_rt/models/groot/pipeline_rtx.py +1179 -0
  293. flash_rt/models/groot/pipeline_rtx_fp16.py +1034 -0
  294. flash_rt/models/groot/pipeline_thor.py +981 -0
  295. flash_rt/models/groot_n17/__init__.py +50 -0
  296. flash_rt/models/groot_n17/calibration.py +439 -0
  297. flash_rt/models/groot_n17/embodiments.py +38 -0
  298. flash_rt/models/groot_n17/mrope_table.py +178 -0
  299. flash_rt/models/groot_n17/pipeline_rtx.py +11 -0
  300. flash_rt/models/groot_n17/pipeline_rtx_fp16.py +836 -0
  301. flash_rt/models/groot_n17/pipeline_rtx_fp8.py +413 -0
  302. flash_rt/models/groot_n17/pipeline_rtx_sm89.py +541 -0
  303. flash_rt/models/groot_n17/pipeline_thor.py +1386 -0
  304. flash_rt/models/higgs_audio_v3/__init__.py +14 -0
  305. flash_rt/models/higgs_audio_v3/_codec/__init__.py +0 -0
  306. flash_rt/models/higgs_audio_v3/_codec/env_guard.py +42 -0
  307. flash_rt/models/higgs_audio_v3/_codec/tokenizer_config.json +129 -0
  308. flash_rt/models/higgs_audio_v3/_codec/tokenizer_model.py +940 -0
  309. flash_rt/models/higgs_audio_v3/codec.py +81 -0
  310. flash_rt/models/higgs_audio_v3/pipeline_rtx.py +64 -0
  311. flash_rt/models/hyvla/__init__.py +1 -0
  312. flash_rt/models/hyvla/pipeline_orin.py +430 -0
  313. flash_rt/models/hyvla/pipeline_thor.py +572 -0
  314. flash_rt/models/lingbot/__init__.py +17 -0
  315. flash_rt/models/lingbot/_csrc_loader.py +70 -0
  316. flash_rt/models/lingbot/buffer_binder.py +156 -0
  317. flash_rt/models/lingbot/calibration.py +163 -0
  318. flash_rt/models/lingbot/forward.py +784 -0
  319. flash_rt/models/lingbot/fp4_ops.py +90 -0
  320. flash_rt/models/lingbot/graph_runner.py +265 -0
  321. flash_rt/models/lingbot/kernel_ops.py +1487 -0
  322. flash_rt/models/lingbot/mixed_attention.py +793 -0
  323. flash_rt/models/lingbot/norms.py +113 -0
  324. flash_rt/models/lingbot/pipeline_thor.py +169 -0
  325. flash_rt/models/lingbot/rope_adapter.py +156 -0
  326. flash_rt/models/lingbot/sample_actions.py +394 -0
  327. flash_rt/models/lingbot/vit.py +486 -0
  328. flash_rt/models/lingbot/vit_rope_adapter.py +247 -0
  329. flash_rt/models/ltx25/__init__.py +16 -0
  330. flash_rt/models/ltx25/_attn_swap.py +244 -0
  331. flash_rt/models/ltx25/_nvfp4_ffn_swap.py +301 -0
  332. flash_rt/models/ltx25/_resident_graph.py +206 -0
  333. flash_rt/models/melband_roformer/__init__.py +13 -0
  334. flash_rt/models/melband_roformer/pipeline.py +329 -0
  335. flash_rt/models/minimax_remover/__init__.py +27 -0
  336. flash_rt/models/minimax_remover/_attention.py +428 -0
  337. flash_rt/models/minimax_remover/_fp8_linear.py +426 -0
  338. flash_rt/models/minimax_remover/_fp8_manual_denoise.py +298 -0
  339. flash_rt/models/minimax_remover/_fp8_pipeline.py +617 -0
  340. flash_rt/models/minimax_remover/_kern_block.py +282 -0
  341. flash_rt/models/minimax_remover/_kernels.py +282 -0
  342. flash_rt/models/minimax_remover/_manual_denoise.py +413 -0
  343. flash_rt/models/minimax_remover/_nvfp4_linear.py +236 -0
  344. flash_rt/models/minimax_remover/_triton_flash_attn.py +139 -0
  345. flash_rt/models/minimax_remover/_utils.py +94 -0
  346. flash_rt/models/minimax_remover/_vae_nvfp4.py +569 -0
  347. flash_rt/models/minimax_remover/_vae_opt.py +880 -0
  348. flash_rt/models/minimax_remover/pipeline.py +209 -0
  349. flash_rt/models/motus/__init__.py +0 -0
  350. flash_rt/models/motus/_action_ffn_v6t_install.py +162 -0
  351. flash_rt/models/motus/_action_und_qkv_fp8_swap.py +266 -0
  352. flash_rt/models/motus/_attn_swap.py +224 -0
  353. flash_rt/models/motus/_awq_fp8_swap.py +525 -0
  354. flash_rt/models/motus/_easycache_swap.py +279 -0
  355. flash_rt/models/motus/_ffn_swap.py +225 -0
  356. flash_rt/models/motus/_fp8_swap.py +635 -0
  357. flash_rt/models/motus/_graph_capture.py +203 -0
  358. flash_rt/models/motus/_handtuned_fp8_dispatch.py +235 -0
  359. flash_rt/models/motus/_kv_cache_swap.py +539 -0
  360. flash_rt/models/motus/_linear_swap.py +248 -0
  361. flash_rt/models/motus/_mixcache_swap.py +276 -0
  362. flash_rt/models/motus/_modulate_fuse_swap.py +1649 -0
  363. flash_rt/models/motus/_motus_nvfp4_ffn_video_swap.py +677 -0
  364. flash_rt/models/motus/_norm_swap.py +240 -0
  365. flash_rt/models/motus/_rope_swap.py +226 -0
  366. flash_rt/models/motus/_stream.py +22 -0
  367. flash_rt/models/motus/_taylorseer_swap.py +275 -0
  368. flash_rt/models/motus/_teacache_swap.py +210 -0
  369. flash_rt/models/motus/_tinyfp8_dispatch_install.py +170 -0
  370. flash_rt/models/motus/_und_ffn_v5t_install.py +180 -0
  371. flash_rt/models/motus/_vae_fp4_swap.py +849 -0
  372. flash_rt/models/motus/_vae_fp8_resample_swap.py +534 -0
  373. flash_rt/models/motus/_vae_fp8_swap.py +1082 -0
  374. flash_rt/models/motus/_vae_swap.py +80 -0
  375. flash_rt/models/motus/_vae_time_conv_fp8_swap.py +292 -0
  376. flash_rt/models/motus/_wan_qkv_fuse_swap.py +905 -0
  377. flash_rt/models/motus/pipeline_rtx.py +1164 -0
  378. flash_rt/models/nexn2/__init__.py +17 -0
  379. flash_rt/models/nexn2/pipeline_rtx.py +137 -0
  380. flash_rt/models/omnivoice/__init__.py +30 -0
  381. flash_rt/models/omnivoice/pipeline_rtx.py +546 -0
  382. flash_rt/models/pi0/__init__.py +9 -0
  383. flash_rt/models/pi0/pipeline_rtx.py +1110 -0
  384. flash_rt/models/pi0/pipeline_thor.py +434 -0
  385. flash_rt/models/pi05/__init__.py +28 -0
  386. flash_rt/models/pi05/pipeline_rtx.py +2209 -0
  387. flash_rt/models/pi05/pipeline_rtx_batched.py +1188 -0
  388. flash_rt/models/pi05/pipeline_rtx_cfg.py +657 -0
  389. flash_rt/models/pi05/pipeline_rtx_cfg_batched.py +435 -0
  390. flash_rt/models/pi05/pipeline_rtx_fp16.py +2276 -0
  391. flash_rt/models/pi05/pipeline_thor.py +929 -0
  392. flash_rt/models/pi05/pipeline_thor_batched.py +346 -0
  393. flash_rt/models/pi05/pipeline_thor_cfg.py +238 -0
  394. flash_rt/models/pi05/pipeline_thor_cfg_batched.py +180 -0
  395. flash_rt/models/pi05/runtime_export.py +449 -0
  396. flash_rt/models/pi0fast/__init__.py +1 -0
  397. flash_rt/models/pi0fast/pipeline.py +840 -0
  398. flash_rt/models/qwen3/__init__.py +11 -0
  399. flash_rt/models/qwen3/pipeline_rtx.py +90 -0
  400. flash_rt/models/qwen36/__init__.py +27 -0
  401. flash_rt/models/qwen36/pipeline_rtx.py +159 -0
  402. flash_rt/models/qwen3_vl/__init__.py +19 -0
  403. flash_rt/models/qwen3_vl/pipeline_rtx.py +145 -0
  404. flash_rt/models/wan22/__init__.py +1 -0
  405. flash_rt/models/wan22/pipeline_rtx.py +28 -0
  406. flash_rt/npu/__init__.py +6 -0
  407. flash_rt/npu/core/__init__.py +0 -0
  408. flash_rt/npu/core/abi.py +81 -0
  409. flash_rt/npu/core/acl_runtime.py +165 -0
  410. flash_rt/npu/core/decode_attention.py +26 -0
  411. flash_rt/npu/core/decoder_int8.py +399 -0
  412. flash_rt/npu/core/device.py +73 -0
  413. flash_rt/npu/core/gu_int8.py +99 -0
  414. flash_rt/npu/core/linear.py +283 -0
  415. flash_rt/npu/core/native_kernels.py +269 -0
  416. flash_rt/npu/core/npu_graph.py +68 -0
  417. flash_rt/npu/frontends/__init__.py +0 -0
  418. flash_rt/npu/frontends/torch/__init__.py +0 -0
  419. flash_rt/npu/frontends/torch/pi05.py +438 -0
  420. flash_rt/npu/hardware/__init__.py +9 -0
  421. flash_rt/npu/models/__init__.py +1 -0
  422. flash_rt/npu/models/pi05/__init__.py +1 -0
  423. flash_rt/npu/models/pi05/attention.py +128 -0
  424. flash_rt/npu/models/pi05/captured.py +212 -0
  425. flash_rt/npu/models/pi05/fast.py +666 -0
  426. flash_rt/npu/models/pi05/pipeline.py +336 -0
  427. flash_rt/npu/models/pi05/quantization.py +235 -0
  428. flash_rt/npu/verify.py +110 -0
  429. flash_rt/py.typed +0 -0
  430. flash_rt/refs/__init__.py +17 -0
  431. flash_rt/refs/pi05_cfg_reference.py +310 -0
  432. flash_rt/runtime/__init__.py +47 -0
  433. flash_rt/runtime/cuda_libraries.py +108 -0
  434. flash_rt/runtime/exec.py +53 -0
  435. flash_rt/runtime/export.py +520 -0
  436. flash_rt/runtime/provider.py +113 -0
  437. flash_rt/runtime/rtc.py +261 -0
  438. flash_rt/runtime/rtc_temporal_fusion.py +545 -0
  439. flash_rt/runtime/vlash.py +420 -0
  440. flash_rt/subgraphs/__init__.py +43 -0
  441. flash_rt/subgraphs/capture.py +179 -0
  442. flash_rt/subgraphs/pi05/__init__.py +3 -0
  443. flash_rt/subgraphs/pi05/context_action.py +67 -0
  444. flash_rt/subgraphs/pi05/rtc_prefix.py +79 -0
  445. flash_rt/subgraphs/pi05/rtc_vjp_guided.py +96 -0
  446. flash_rt/subgraphs/pi05/stage_plans.py +89 -0
  447. flash_rt/subgraphs/pi05/vlash.py +29 -0
  448. flash_rt/subgraphs/stage_plan.py +214 -0
  449. flash_rt/utils/__init__.py +1 -0
  450. flash_rt/utils/paligemma_tokenizer.py +135 -0
  451. flash_rt-0.2.0.dist-info/METADATA +1341 -0
  452. flash_rt-0.2.0.dist-info/RECORD +455 -0
  453. flash_rt-0.2.0.dist-info/WHEEL +5 -0
  454. flash_rt-0.2.0.dist-info/licenses/LICENSE +202 -0
  455. flash_rt-0.2.0.dist-info/top_level.txt +1 -0
@@ -0,0 +1,549 @@
1
+ #!/usr/bin/env python3
2
+ """Cosmos3-Edge AV inverse-dynamics denoise compute path (Thor SM110).
3
+
4
+ Per docs/adding_new_model.md §0 rule 1 this is the model's own
5
+ ``models/cosmos3_edge/pipeline_thor.py``: the quantized two-tower MoT
6
+ kernel-call sequence, the static und K/V cache, and the whole-denoise
7
+ CUDA-graph capture live here, self-contained.
8
+
9
+ The gen tower carries static vision conditioning plus 60 action tokens; the
10
+ head is the per-domain action projection. The und (text) tower is identical
11
+ across denoise steps, so its per-layer K/V (k_norm_und_for_gen + rope applied)
12
+ is computed once and cached; each step runs only the gen tower against the
13
+ cached und K/V through FA4 attention.
14
+
15
+ Precision (selected by the caller, not the environment):
16
+ - quant="fp8" : w8a8 FP8 E4M3 GEMMs (fp8_gemm_descale_bf16out) with
17
+ per-tensor weight scales and device-side dynamic activation
18
+ scales. Production default.
19
+ - quant="bf16" : reference-accuracy path (bf16 GemmRunner GEMMs).
20
+ bf16_projs keeps named projections ("q", "k", "v", "o", "up", "down") in bf16.
21
+ This module reads no environment variables.
22
+ """
23
+ from __future__ import annotations
24
+
25
+ import torch
26
+
27
+ import flash_rt.flash_rt_kernels as fvk
28
+ from flash_rt.frontends.torch._cosmos3_edge_thor_spec import SPEC
29
+ from flash_rt.models.cosmos3_edge.boundary_dump import EdgeBoundaryDump
30
+ from flash_rt.models.cosmos3_edge.dump_replay import EDGE_ACTION_MODEL_SHAPE, EDGE_FLAT_DIM
31
+ from flash_rt.models.cosmos3_edge.static_engine import EdgeStaticBufferEngine
32
+ from flash_rt.models.cosmos3_edge.static_unipc import EdgeStaticUniPCScheduler
33
+ from flash_rt.models.cosmos3_edge.weights import EdgeTransformerWeights
34
+
35
+ DEV = "cuda"
36
+ BF = torch.bfloat16
37
+ H, KVH, D, FF, HID, NL = (
38
+ SPEC.num_heads,
39
+ SPEC.num_kv_heads,
40
+ SPEC.head_dim,
41
+ SPEC.ffn_size,
42
+ SPEC.hidden_size,
43
+ SPEC.num_layers,
44
+ )
45
+ EPS = SPEC.rms_eps
46
+ NA, AD = EDGE_ACTION_MODEL_SHAPE
47
+ ALL_PROJS = ("q", "k", "v", "o", "up", "down")
48
+
49
+ _GEN_PROJ_KEYS = {
50
+ "q": "self_attn.add_q_proj.weight",
51
+ "k": "self_attn.add_k_proj.weight",
52
+ "v": "self_attn.add_v_proj.weight",
53
+ "o": "self_attn.to_add_out.weight",
54
+ "up": "mlp_moe_gen.up_proj.weight",
55
+ "down": "mlp_moe_gen.down_proj.weight",
56
+ }
57
+
58
+
59
+ class CosmosEdgeThor:
60
+ """Static-buffer quantized gen-tower denoise engine for Cosmos3-Edge AV."""
61
+
62
+ def __init__(
63
+ self,
64
+ weights: EdgeTransformerWeights,
65
+ boundary_dump: EdgeBoundaryDump,
66
+ timesteps: tuple[int, ...],
67
+ *,
68
+ quant: str = "fp8",
69
+ bf16_projs: tuple[str, ...] = (),
70
+ qkv_fused: bool = False,
71
+ ffn_fp4: bool = False,
72
+ slim_last: bool = True,
73
+ shift: float,
74
+ ):
75
+ self.slim_last = bool(slim_last)
76
+ if quant not in ("fp8", "bf16"):
77
+ raise ValueError(f"unsupported quant mode: {quant}")
78
+ self.quant = quant
79
+ self.fp4 = None
80
+ self.ffn_fp4 = False
81
+ if ffn_fp4 and quant == "fp8" and not ({"up", "down"} & set(bf16_projs)):
82
+ import flash_rt.flash_rt_fp4 as fp4mod
83
+
84
+ if not fp4mod.has_nvfp4():
85
+ raise RuntimeError("ffn_fp4 requested but flash_rt_fp4 lacks NVFP4 support")
86
+ self.fp4 = fp4mod
87
+ self.ffn_fp4 = True
88
+ self.bf16_projs = set(ALL_PROJS) if quant == "bf16" else {p for p in bf16_projs if p}
89
+ self.num_steps = len(timesteps)
90
+
91
+ # Static conditioning, und K/V cache, action-head constants, and the
92
+ # timestep-embed table come from the validated bring-up engine; the
93
+ # denoise hot path below never calls back into it.
94
+ base = EdgeStaticBufferEngine(weights, boundary_dump, device=DEV, dtype=BF)
95
+ base.precompute_timestep_embeds(timesteps)
96
+ fa4_getter = getattr(base.reference, "_get_fa4_fwd", None)
97
+ self._fa4_fwd = fa4_getter() if fa4_getter is not None else None
98
+ if self._fa4_fwd is None:
99
+ raise RuntimeError("CosmosEdgeThor requires the FA4 forward entry on this build")
100
+
101
+ self.NU = int(base.und_cache[0][0].shape[0])
102
+ self.NG = int(base.full.shape[0])
103
+ self.NJ = self.NU + self.NG
104
+ NU, NG, NJ = self.NU, self.NG, self.NJ
105
+
106
+ self.gemm = fvk.GemmRunner()
107
+ z = lambda *s: torch.zeros(*s, device=DEV, dtype=BF)
108
+
109
+ # --- weights: quantized gen-tower projections + bf16 norms ---
110
+ def qf8(w_nk: torch.Tensor): # per-tensor FP8 E4M3; GEMM wants B as [K,N]
111
+ w = w_nk.t().contiguous()
112
+ s = max(w.float().abs().max().item() / 448.0, 1e-12)
113
+ f8 = (w.float() / s).clamp(-448, 448).to(torch.float8_e4m3fn).contiguous()
114
+ return f8, torch.tensor([s], dtype=torch.float32, device=DEV)
115
+
116
+ # QKV fused wide GEMM measured NEGATIVE on Thor (P50 6.056s vs 5.883s
117
+ # separate: cuBLASLt N=4096 tactic worse than 3×separate + strided V
118
+ # copy). Keep the path for re-evaluation; off by default.
119
+ self.qkv_fused = bool(qkv_fused) and not ({"q", "k", "v"} & self.bf16_projs) and quant == "fp8"
120
+ self.Wt = {}
121
+ self.Wf8 = {}
122
+ self.Wds = {}
123
+ self.Wf4 = {}
124
+ self.Wn = {}
125
+ for li in range(NL):
126
+ loaded = {}
127
+ for nm, key in _GEN_PROJ_KEYS.items():
128
+ loaded[nm] = weights.load_tensor(f"layers.{li}.{key}", device=DEV, dtype=BF)
129
+ if self.qkv_fused:
130
+ merged = torch.cat([loaded.pop("q"), loaded.pop("k"), loaded.pop("v")], dim=0)
131
+ self.Wf8[(li, "qkv")], self.Wds[(li, "qkv")] = qf8(merged)
132
+ del merged
133
+ if self.ffn_fp4:
134
+ for nm in ("up", "down"):
135
+ w16 = loaded.pop(nm).to(torch.float16).contiguous()
136
+ n_out, k_in = w16.shape
137
+ packed = torch.empty(n_out, k_in // 2, dtype=torch.uint8, device=DEV)
138
+ sfb = torch.empty(
139
+ self.fp4.sfa_size_bytes(n_out, k_in, True), dtype=torch.uint8, device=DEV)
140
+ rc = self.fp4.quantize_fp4_dynamic_sfa_mse_fp16(
141
+ w16.data_ptr(), packed.data_ptr(), sfb.data_ptr(), n_out, k_in, True, 0)
142
+ if rc != 0:
143
+ raise RuntimeError(f"fp4 weight quant failed rc={rc} layer {li} {nm}")
144
+ self.Wf4[(li, nm)] = (packed, sfb)
145
+ del w16
146
+ for nm, w in loaded.items():
147
+ if nm in self.bf16_projs:
148
+ self.Wt[(li, nm)] = w.t().contiguous()
149
+ else:
150
+ self.Wf8[(li, nm)], self.Wds[(li, nm)] = qf8(w)
151
+ for nm, key in (
152
+ ("in_ln", "input_layernorm_moe_gen.weight"),
153
+ ("post_ln", "post_attention_layernorm_moe_gen.weight"),
154
+ ("q_norm", "self_attn.norm_added_q.weight"),
155
+ ("k_norm", "self_attn.norm_added_k.weight"),
156
+ ):
157
+ self.Wn[(li, nm)] = weights.load_tensor(f"layers.{li}.{key}", device=DEV, dtype=BF)
158
+ self.norm_g = weights.load_tensor("norm_moe_gen.weight", device=DEV, dtype=BF)
159
+
160
+ # --- action head constants (already domain-selected by the base engine) ---
161
+ self.action_in_w = base.action_in_w
162
+ self.action_static_bias = base.action_static_bias
163
+ self.action_out_w = base.action_out_w
164
+ self.action_out_bias = base.action_out_bias
165
+ self.raw_action_dim = base.raw_action_dim
166
+ self.action_indexes = base.action_indexes.to(torch.int64).contiguous()
167
+ self.t_emb = base.timestep_embed_cache
168
+ if self.t_emb is None or self.t_emb.shape[0] != self.num_steps:
169
+ raise RuntimeError("timestep embed table missing or mismatched")
170
+
171
+ # --- static activation buffers ---
172
+ self.gen_init = base.full.clone() # vision conditioning; action rows zero
173
+ self.Hb = z(NG, HID)
174
+ self.nrm = z(NG, HID)
175
+ self.nrm2 = z(NG, HID)
176
+ self.Qt = z(NG, H * D)
177
+ self.Kt = z(NG, KVH * D)
178
+ self.QKVt = z(NG, (H + 2 * KVH) * D) if self.qkv_fused else None
179
+ self.Qb = z(NG, H, D)
180
+ self.attn = z(1, NG, H, D)
181
+ self.lse = torch.empty(1, H, NG, dtype=torch.float32, device=DEV)
182
+ self.ob = z(NG, HID)
183
+ self.up = z(NG, FF)
184
+ self.dn = z(NG, HID)
185
+ self.action_input = z(NA, AD)
186
+ self.action_encoded = z(NA, HID)
187
+ self.action_hidden = z(NA, HID)
188
+ self.action_norm = z(NA, HID)
189
+ self.action_out = z(NA, AD)
190
+ # Slim last layer: only the NA action rows feed the head, so the final
191
+ # layer's Q/attention/o/FFN run at M=NA (K/V still need all rows).
192
+ self.Qs = z(NA, H, D)
193
+ self.attn_s = z(1, NA, H, D)
194
+ self.lse_s = torch.empty(1, H, NA, dtype=torch.float32, device=DEV)
195
+ self.velocity = z(EDGE_FLAT_DIM)
196
+ self.latent = torch.zeros(EDGE_FLAT_DIM, device=DEV, dtype=torch.float32)
197
+ if self.ffn_fp4:
198
+ self.a4h = torch.empty(NG, HID // 2, dtype=torch.uint8, device=DEV)
199
+ self.sfa_h = torch.empty(self.fp4.sfa_size_bytes(NG, HID, False), dtype=torch.uint8, device=DEV)
200
+ self.up16 = torch.empty(NG, FF, dtype=torch.float16, device=DEV)
201
+ self.a4f = torch.empty(NG, FF // 2, dtype=torch.uint8, device=DEV)
202
+ self.sfa_f = torch.empty(self.fp4.sfa_size_bytes(NG, FF, False), dtype=torch.uint8, device=DEV)
203
+ self.dn16 = torch.empty(NG, HID, dtype=torch.float16, device=DEV)
204
+ # Fastest verified GEMM variant at the Edge FFN shapes; -1 = default API.
205
+ self.fp4_variant = 3
206
+ probe = torch.zeros(2, HID // 2, dtype=torch.uint8, device=DEV)
207
+ probe_sfa = torch.zeros(self.fp4.sfa_size_bytes(2, HID, False), dtype=torch.uint8, device=DEV)
208
+ probe_out = torch.empty(2, FF, dtype=torch.float16, device=DEV)
209
+ rc = self.fp4.cutlass_fp4_gemm_variant(
210
+ self.fp4_variant, probe.data_ptr(), probe_sfa.data_ptr(),
211
+ self.Wf4[(0, "up")][0].data_ptr(), self.Wf4[(0, "up")][1].data_ptr(),
212
+ probe_out.data_ptr(), 2, FF, HID, 1.0, 0.0, 0)
213
+ if rc != 0:
214
+ self.fp4_variant = -1
215
+ probe_relu = torch.empty(2, FF // 2, dtype=torch.uint8, device=DEV)
216
+ probe_relu_sfa = torch.empty(
217
+ self.fp4.sfa_size_bytes(2, FF, False), dtype=torch.uint8, device=DEV)
218
+ rc = self.fp4.cosmos3_edge_fp4_gemm_relu2_fp4out(
219
+ probe.data_ptr(), probe_sfa.data_ptr(),
220
+ self.Wf4[(0, "up")][0].data_ptr(), self.Wf4[(0, "up")][1].data_ptr(),
221
+ probe_relu.data_ptr(), probe_relu_sfa.data_ptr(), 2, FF, HID, 0)
222
+ self.fp4_relu2_epilogue = rc == 0
223
+ del probe, probe_sfa, probe_out, probe_relu, probe_relu_sfa
224
+ self.a8h = self.a8f = self.asc = None
225
+ # Per-(layer, site) static activation scales for the fused quant chain.
226
+ # Sites: 0 = qkv input, 1 = attention output, 2 = FFN input, 3 = down input.
227
+ self.site_scale = torch.full((NL, 4), 1e-12, dtype=torch.float32, device=DEV)
228
+ self.calibrated = False
229
+ if quant == "fp8":
230
+ self.a8h = torch.empty(NG, HID, dtype=torch.float8_e4m3fn, device=DEV)
231
+ self.a8f = torch.empty(NG, FF, dtype=torch.float8_e4m3fn, device=DEV)
232
+ self.asc = torch.empty(1, dtype=torch.float32, device=DEV)
233
+
234
+ # --- per-layer joint K/V with the static und prefix pre-installed ---
235
+ self.Kj = torch.zeros(NL, 1, NJ, KVH, D, device=DEV, dtype=BF)
236
+ self.Vj = torch.zeros(NL, 1, NJ, KVH, D, device=DEV, dtype=BF)
237
+ for li, (k_und, v_und) in enumerate(base.und_cache):
238
+ self.Kj[li, 0, :NU].copy_(k_und)
239
+ self.Vj[li, 0, :NU].copy_(v_und)
240
+
241
+ # --- static rope tables (full_only sequence) ---
242
+ t = boundary_dump.tensors
243
+ self.cos_f = t["s00/layers/00/rope/cos/full_only_seq"].to(device=DEV, dtype=BF).contiguous()
244
+ self.sin_f = t["s00/layers/00/rope/sin/full_only_seq"].to(device=DEV, dtype=BF).contiguous()
245
+
246
+ self.compute_steps: frozenset[int] | None = None
247
+ self.unipc = EdgeStaticUniPCScheduler(self.num_steps, device=torch.device(DEV), shift=shift)
248
+ if not self.unipc.native_available:
249
+ raise RuntimeError("native UniPC step binding is required")
250
+ self.graph = None
251
+ del base
252
+ torch.cuda.synchronize()
253
+
254
+ # ---- kernel helpers ----
255
+ def _s(self):
256
+ return torch.cuda.current_stream().cuda_stream
257
+
258
+ def _site_ptr(self, li: int, site: int) -> int:
259
+ return self.site_scale.data_ptr() + 4 * (li * 4 + site)
260
+
261
+ def _proj(
262
+ self,
263
+ a_bf16: torch.Tensor,
264
+ a_f8: torch.Tensor | None,
265
+ li: int,
266
+ nm: str,
267
+ out: torch.Tensor,
268
+ n_out: int,
269
+ act_scale_ptr: int | None = None,
270
+ ):
271
+ m, k = a_bf16.shape
272
+ if (li, nm) in self.Wf8:
273
+ assert a_f8 is not None
274
+ scale_ptr = self.asc.data_ptr() if act_scale_ptr is None else act_scale_ptr
275
+ fvk.fp8_gemm_descale_bf16out(
276
+ a_f8.data_ptr(), self.Wf8[(li, nm)].data_ptr(), out.data_ptr(),
277
+ m, n_out, k, scale_ptr, self.Wds[(li, nm)].data_ptr(), self._s(),
278
+ )
279
+ else:
280
+ self.gemm.bf16_nn(a_bf16.data_ptr(), self.Wt[(li, nm)].data_ptr(), out.data_ptr(), m, n_out, k, self._s())
281
+
282
+ def _fp4_gemm(self, a_packed, a_sfa, li: int, nm: str, out16, m: int, n: int, k: int, s: int):
283
+ w_packed, w_sfb = self.Wf4[(li, nm)]
284
+ if self.fp4_variant >= 0:
285
+ rc = self.fp4.cutlass_fp4_gemm_variant(
286
+ self.fp4_variant, a_packed.data_ptr(), a_sfa.data_ptr(),
287
+ w_packed.data_ptr(), w_sfb.data_ptr(), out16.data_ptr(), m, n, k, 1.0, 0.0, s)
288
+ else:
289
+ rc = self.fp4.cutlass_fp4_sq_fp16(
290
+ a_packed.data_ptr(), a_sfa.data_ptr(),
291
+ w_packed.data_ptr(), w_sfb.data_ptr(), out16.data_ptr(), m, n, k, 1.0, 0.0, s)
292
+ if rc != 0:
293
+ raise RuntimeError(f"fp4 GEMM failed rc={rc} layer {li} {nm}")
294
+
295
+ def _ffn_fp4_block(self, li: int, x_bf16: torch.Tensor, s: int):
296
+ """FP4 W4A4 FFN: fused bf16 res+rms -> fp4, up GEMM, relu2 -> fp4, down GEMM."""
297
+ NG = self.NG
298
+ self.fp4.cosmos3_edge_res_rms_fp4_sfa_bf16(
299
+ self.Hb.data_ptr(), x_bf16.data_ptr(), self.Wn[(li, "post_ln")].data_ptr(),
300
+ self.a4h.data_ptr(), self.sfa_h.data_ptr(), NG, HID, EPS, s)
301
+ if self.fp4_relu2_epilogue:
302
+ w_packed, w_sfb = self.Wf4[(li, "up")]
303
+ rc = self.fp4.cosmos3_edge_fp4_gemm_relu2_fp4out(
304
+ self.a4h.data_ptr(), self.sfa_h.data_ptr(),
305
+ w_packed.data_ptr(), w_sfb.data_ptr(),
306
+ self.a4f.data_ptr(), self.sfa_f.data_ptr(), NG, FF, HID, s)
307
+ if rc != 0:
308
+ raise RuntimeError(f"fused FP4 up+relu2 failed rc={rc} layer {li}")
309
+ else:
310
+ self._fp4_gemm(self.a4h, self.sfa_h, li, "up", self.up16, NG, FF, HID, s)
311
+ self.fp4.cosmos3_edge_relu2_fp4_sfa_fp16(
312
+ self.up16.data_ptr(), self.a4f.data_ptr(), self.sfa_f.data_ptr(), NG, FF, s)
313
+ self._fp4_gemm(self.a4f, self.sfa_f, li, "down", self.dn16, NG, HID, FF, s)
314
+ self.dn.copy_(self.dn16)
315
+
316
+ def _quant(self, a_bf16: torch.Tensor, a_f8: torch.Tensor, needed: bool, li: int, site: int):
317
+ """Dynamic quantization pass; accumulates the per-site scale ceiling."""
318
+ if needed and self.quant == "fp8":
319
+ fvk.quantize_fp8_device(a_bf16.data_ptr(), a_f8.data_ptr(), self.asc.data_ptr(), a_bf16.numel(), self._s())
320
+ fvk.fp8_accumulate_scale_max(self.asc.data_ptr(), self._site_ptr(li, site), self._s())
321
+
322
+ def _layer_needs_fp8(self, li: int, names: tuple[str, ...]) -> bool:
323
+ return any((li, nm) in self.Wf8 for nm in names)
324
+
325
+ @property
326
+ def _fused_static(self) -> bool:
327
+ return self.quant == "fp8" and self.calibrated and not self.bf16_projs
328
+
329
+ # ---- per-step forward: latent(f32 flat) -> velocity(bf16 flat) ----
330
+ def forward_step(self, step: int, latent: torch.Tensor) -> torch.Tensor:
331
+ s = self._s()
332
+ NG, NJ, NU = self.NG, self.NJ, self.NU
333
+
334
+ fvk.cosmos3_edge_copy_action_tail_f32_to_bf16(
335
+ latent.data_ptr(), self.action_input.data_ptr(), latent.numel(), NA * AD, s)
336
+ self.gemm.bf16_nn(
337
+ self.action_input.data_ptr(), self.action_in_w.data_ptr(),
338
+ self.action_encoded.data_ptr(), NA, HID, AD, s)
339
+ fvk.cosmos3_edge_add_action_bias_timestep_bf16(
340
+ self.action_encoded.data_ptr(), self.action_static_bias.data_ptr(),
341
+ self.t_emb[step : step + 1].data_ptr(), NA, HID, s)
342
+ self.Hb.copy_(self.gen_init)
343
+ fvk.cosmos3_edge_scatter_rows_bf16(
344
+ self.action_encoded.data_ptr(), self.Hb.data_ptr(), self.action_indexes.data_ptr(), NA, HID, s)
345
+
346
+ fused = self._fused_static
347
+ if fused:
348
+ fvk.rms_norm(self.Hb.data_ptr(), self.Wn[(0, "in_ln")].data_ptr(), self.nrm.data_ptr(), NG, HID, EPS, s)
349
+ fvk.quantize_fp8_static(self.nrm.data_ptr(), self.a8h.data_ptr(), self._site_ptr(0, 0), NG * HID, s)
350
+ else:
351
+ fvk.rms_norm(self.Hb.data_ptr(), self.Wn[(0, "in_ln")].data_ptr(), self.nrm.data_ptr(), NG, HID, EPS, s)
352
+ slim_done = False
353
+ for li in range(NL):
354
+ qkv_scale = self._site_ptr(li, 0) if fused else None
355
+ if not fused:
356
+ self._quant(
357
+ self.nrm, self.a8h,
358
+ self.qkv_fused or self._layer_needs_fp8(li, ("q", "k", "v")), li, 0)
359
+ if self.qkv_fused:
360
+ wide = (H + 2 * KVH) * D
361
+ self._proj(self.nrm, self.a8h, li, "qkv", self.QKVt, wide, qkv_scale)
362
+ fvk.cosmos3_edge_qk_norm_rope_strided_bf16(
363
+ self.QKVt.data_ptr(), self.QKVt.data_ptr() + 2 * H * D,
364
+ self.Wn[(li, "q_norm")].data_ptr(), self.Wn[(li, "k_norm")].data_ptr(),
365
+ self.cos_f.data_ptr(), self.sin_f.data_ptr(),
366
+ self.Qb.data_ptr(), self.Kj[li, 0, NU:NJ].data_ptr(),
367
+ NG, H, KVH, wide, wide, EPS, s)
368
+ self.Vj[li, 0, NU:NJ].view(NG, KVH * D).copy_(self.QKVt[:, (H + KVH) * D :])
369
+ else:
370
+ self._proj(self.nrm, self.a8h, li, "q", self.Qt, H * D, qkv_scale)
371
+ self._proj(self.nrm, self.a8h, li, "k", self.Kt, KVH * D, qkv_scale)
372
+ self._proj(self.nrm, self.a8h, li, "v", self.Vj[li, 0, NU:NJ].view(NG, KVH * D), KVH * D, qkv_scale)
373
+ fvk.cosmos3_edge_qk_norm_rope_bf16(
374
+ self.Qt.data_ptr(), self.Kt.data_ptr(),
375
+ self.Wn[(li, "q_norm")].data_ptr(), self.Wn[(li, "k_norm")].data_ptr(),
376
+ self.cos_f.data_ptr(), self.sin_f.data_ptr(),
377
+ self.Qb.data_ptr(), self.Kj[li, 0, NU:NJ].data_ptr(),
378
+ NG, H, KVH, D, D, EPS, s)
379
+ if fused and self.slim_last and li == NL - 1:
380
+ # Last layer: only the NA action rows are consumed by the head.
381
+ # K/V (computed above) still cover all rows; Q/attention/o/FFN
382
+ # shrink to M=NA. Site scales calibrated at full M are upper
383
+ # bounds for the row subset.
384
+ fvk.cosmos3_edge_gather_rows_bf16(
385
+ self.Qb.data_ptr(), self.Qs.data_ptr(), self.action_indexes.data_ptr(), NA, HID, s)
386
+ self._fa4_fwd(
387
+ self.Qs.view(1, NA, H, D), self.Kj[li], self.Vj[li],
388
+ softmax_scale=D ** -0.5, causal=False, pack_gqa=True,
389
+ out=self.attn_s, lse=self.lse_s)
390
+ attn_s2d = self.attn_s.view(NA, HID)
391
+ fvk.quantize_fp8_static(
392
+ attn_s2d.data_ptr(), self.a8h.data_ptr(), self._site_ptr(li, 1), NA * HID, s)
393
+ self._proj(attn_s2d, self.a8h, li, "o", self.ob, HID, self._site_ptr(li, 1))
394
+ fvk.cosmos3_edge_gather_rows_bf16(
395
+ self.Hb.data_ptr(), self.action_hidden.data_ptr(), self.action_indexes.data_ptr(), NA, HID, s)
396
+ if self.ffn_fp4:
397
+ self.fp4.cosmos3_edge_res_rms_fp4_sfa_bf16(
398
+ self.action_hidden.data_ptr(), self.ob.data_ptr(),
399
+ self.Wn[(li, "post_ln")].data_ptr(),
400
+ self.a4h.data_ptr(), self.sfa_h.data_ptr(), NA, HID, EPS, s)
401
+ if self.fp4_relu2_epilogue:
402
+ w_packed, w_sfb = self.Wf4[(li, "up")]
403
+ rc = self.fp4.cosmos3_edge_fp4_gemm_relu2_fp4out(
404
+ self.a4h.data_ptr(), self.sfa_h.data_ptr(),
405
+ w_packed.data_ptr(), w_sfb.data_ptr(),
406
+ self.a4f.data_ptr(), self.sfa_f.data_ptr(), NA, FF, HID, s)
407
+ if rc != 0:
408
+ raise RuntimeError(f"fused FP4 up+relu2 failed rc={rc} layer {li}")
409
+ else:
410
+ self._fp4_gemm(self.a4h, self.sfa_h, li, "up", self.up16, NA, FF, HID, s)
411
+ self.fp4.cosmos3_edge_relu2_fp4_sfa_fp16(
412
+ self.up16.data_ptr(), self.a4f.data_ptr(), self.sfa_f.data_ptr(), NA, FF, s)
413
+ self._fp4_gemm(self.a4f, self.sfa_f, li, "down", self.dn16, NA, HID, FF, s)
414
+ self.dn[:NA].copy_(self.dn16[:NA])
415
+ else:
416
+ fvk.residual_add_rms_norm_fp8(
417
+ self.action_hidden.data_ptr(), self.ob.data_ptr(),
418
+ self.Wn[(li, "post_ln")].data_ptr(),
419
+ self.a8h.data_ptr(), NA, HID, EPS, self._site_ptr(li, 2), s)
420
+ self._proj(self.nrm2[:NA], self.a8h, li, "up", self.up, FF, self._site_ptr(li, 2))
421
+ fvk.cosmos3_edge_relu2_to_fp8_static_bf16(
422
+ self.up.data_ptr(), self.a8f.data_ptr(), self._site_ptr(li, 3), NA * FF, s)
423
+ self._proj(self.up[:NA], self.a8f, li, "down", self.dn, HID, self._site_ptr(li, 3))
424
+ fvk.residual_add(self.action_hidden.data_ptr(), self.dn.data_ptr(), NA * HID, s)
425
+ slim_done = True
426
+ break
427
+ self._fa4_fwd(
428
+ self.Qb.view(1, NG, H, D), self.Kj[li], self.Vj[li],
429
+ softmax_scale=D ** -0.5, causal=False, pack_gqa=True,
430
+ out=self.attn, lse=self.lse)
431
+ attn2d = self.attn.view(NG, HID)
432
+ if fused:
433
+ fvk.quantize_fp8_static(attn2d.data_ptr(), self.a8h.data_ptr(), self._site_ptr(li, 1), NG * HID, s)
434
+ self._proj(attn2d, self.a8h, li, "o", self.ob, HID, self._site_ptr(li, 1))
435
+ if self.ffn_fp4:
436
+ self._ffn_fp4_block(li, self.ob, s)
437
+ else:
438
+ fvk.residual_add_rms_norm_fp8(
439
+ self.Hb.data_ptr(), self.ob.data_ptr(), self.Wn[(li, "post_ln")].data_ptr(),
440
+ self.a8h.data_ptr(), NG, HID, EPS, self._site_ptr(li, 2), s)
441
+ self._proj(self.nrm2, self.a8h, li, "up", self.up, FF, self._site_ptr(li, 2))
442
+ fvk.cosmos3_edge_relu2_to_fp8_static_bf16(
443
+ self.up.data_ptr(), self.a8f.data_ptr(), self._site_ptr(li, 3), NG * FF, s)
444
+ self._proj(self.up, self.a8f, li, "down", self.dn, HID, self._site_ptr(li, 3))
445
+ if li + 1 < NL:
446
+ fvk.residual_add_rms_norm_fp8(
447
+ self.Hb.data_ptr(), self.dn.data_ptr(), self.Wn[(li + 1, "in_ln")].data_ptr(),
448
+ self.a8h.data_ptr(), NG, HID, EPS, self._site_ptr(li + 1, 0), s)
449
+ else:
450
+ fvk.residual_add(self.Hb.data_ptr(), self.dn.data_ptr(), NG * HID, s)
451
+ else:
452
+ self._quant(attn2d, self.a8h, self._layer_needs_fp8(li, ("o",)), li, 1)
453
+ self._proj(attn2d, self.a8h, li, "o", self.ob, HID)
454
+ if self.ffn_fp4:
455
+ self._ffn_fp4_block(li, self.ob, s)
456
+ else:
457
+ fvk.residual_add_rms_norm(
458
+ self.Hb.data_ptr(), self.ob.data_ptr(), self.Wn[(li, "post_ln")].data_ptr(),
459
+ self.nrm2.data_ptr(), NG, HID, EPS, s)
460
+ self._quant(self.nrm2, self.a8h, self._layer_needs_fp8(li, ("up",)), li, 2)
461
+ self._proj(self.nrm2, self.a8h, li, "up", self.up, FF)
462
+ fvk.relu2_inplace_bf16(self.up.data_ptr(), NG * FF, s)
463
+ self._quant(self.up, self.a8f, self._layer_needs_fp8(li, ("down",)), li, 3)
464
+ self._proj(self.up, self.a8f, li, "down", self.dn, HID)
465
+ if li + 1 < NL:
466
+ fvk.residual_add_rms_norm(
467
+ self.Hb.data_ptr(), self.dn.data_ptr(), self.Wn[(li + 1, "in_ln")].data_ptr(),
468
+ self.nrm.data_ptr(), NG, HID, EPS, s)
469
+ else:
470
+ fvk.residual_add(self.Hb.data_ptr(), self.dn.data_ptr(), NG * HID, s)
471
+
472
+ if not slim_done:
473
+ fvk.cosmos3_edge_gather_rows_bf16(
474
+ self.Hb.data_ptr(), self.action_hidden.data_ptr(), self.action_indexes.data_ptr(), NA, HID, s)
475
+ fvk.rms_norm(
476
+ self.action_hidden.data_ptr(), self.norm_g.data_ptr(), self.action_norm.data_ptr(), NA, HID, EPS, s)
477
+ self.gemm.bf16_nn(
478
+ self.action_norm.data_ptr(), self.action_out_w.data_ptr(), self.action_out.data_ptr(), NA, AD, HID, s)
479
+ fvk.cosmos3_edge_add_bias_zero_action_tail_bf16(
480
+ self.action_out.data_ptr(), self.action_out_bias.data_ptr(), NA, AD, self.raw_action_dim, s)
481
+ fvk.cosmos3_edge_fill_flat_velocity_bf16(
482
+ self.action_out.data_ptr(), self.velocity.data_ptr(), self.velocity.numel(), NA * AD, s)
483
+ return self.velocity
484
+
485
+ def set_teacache(self, compute_steps: tuple[int, ...] | list[int] | None) -> None:
486
+ """Fixed compute-step schedule: skipped steps reuse the last computed
487
+ velocity (the static velocity buffer) while UniPC still advances every
488
+ step. Step 0 must always compute. Call before ``capture``."""
489
+ if self.graph is not None:
490
+ raise RuntimeError("set_teacache must be called before graph capture")
491
+ if compute_steps is None:
492
+ self.compute_steps = None
493
+ return
494
+ steps = sorted(set(int(v) for v in compute_steps))
495
+ if not steps or steps[0] != 0 or steps[-1] >= self.num_steps:
496
+ raise ValueError(f"invalid TeaCache compute schedule: {steps}")
497
+ self.compute_steps = frozenset(steps)
498
+
499
+ # ---- whole-denoise loop: forward + native UniPC per step ----
500
+ def run_loop(self) -> torch.Tensor:
501
+ # UniPC state buffers are allocated once; their contents are fully
502
+ # rewritten in step order on every run (step 0 reads no history).
503
+ if self.unipc.prev_m1 is None:
504
+ self.unipc.reset(self.latent)
505
+ for step in range(self.num_steps):
506
+ if self.compute_steps is None or step in self.compute_steps:
507
+ velocity = self.forward_step(step, self.latent)
508
+ else:
509
+ velocity = self.velocity
510
+ self.unipc.step(self.latent, velocity, step)
511
+ return self.latent
512
+
513
+ def calibrate(self, noise_flat_f32: torch.Tensor, *, margin: float = 1.25) -> None:
514
+ """One dynamic-quant denoise pass to record per-site scale ceilings.
515
+
516
+ The recorded ceilings (times ``margin``) become the static scales for
517
+ the fused quant chain; FP8 E4M3 saturates gracefully on the rare
518
+ excursion past the ceiling.
519
+ """
520
+ if self.quant != "fp8" or self.bf16_projs:
521
+ return
522
+ self.calibrated = False
523
+ self.latent.copy_(noise_flat_f32.to(device=DEV, dtype=torch.float32))
524
+ self.run_loop()
525
+ torch.cuda.synchronize()
526
+ self.site_scale.mul_(margin)
527
+ self.calibrated = True
528
+
529
+ def capture(self, warmup_noise: torch.Tensor | None = None) -> None:
530
+ if warmup_noise is not None:
531
+ self.latent.copy_(warmup_noise.to(device=DEV, dtype=torch.float32))
532
+ st = torch.cuda.Stream()
533
+ st.wait_stream(torch.cuda.current_stream())
534
+ with torch.cuda.stream(st):
535
+ for _ in range(2):
536
+ self.run_loop()
537
+ torch.cuda.current_stream().wait_stream(st)
538
+ self.graph = torch.cuda.CUDAGraph()
539
+ with torch.cuda.graph(self.graph):
540
+ self.run_loop()
541
+
542
+ def denoise(self, noise_flat_f32: torch.Tensor) -> torch.Tensor:
543
+ """Full 30-step denoise from flat noise; returns the final flat latent."""
544
+ self.latent.copy_(noise_flat_f32.to(device=DEV, dtype=torch.float32))
545
+ if self.graph is not None:
546
+ self.graph.replay()
547
+ else:
548
+ self.run_loop()
549
+ return self.latent