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,1013 @@
1
+ """FlashRT AMD -- GROOT N1.7 FP8 torch frontend for CDNA4 (MI350X, gfx950).
2
+
3
+ FP8 kernel backbone with an unquantized bf16 action head. The whole
4
+ VLM backbone (ViT / DeepStack / LLM / VL self-attn) runs through the AMD
5
+ FP8 kernel surface via :mod:`flash_rt.amd.models.groot_n17.pipeline`
6
+ (FUSED-EPILOGUE tier by default — bias / bias+GELU in the hipBLASLt FP8
7
+ epilogue with host alphas; ``FVK_AMD_FUSED_EPILOGUE=0`` falls back to the
8
+ decomposed descale form), and the DiT action head stays bf16
9
+ (``_DIT_USE_FP8 = False``). The same env picks the DiT driver: fused on →
10
+ the AMD-local ``dit_forward`` (``bf16_nn_bias`` / ``bf16_nn_bias_gelu``);
11
+ off → the hardware-independent
12
+ :func:`flash_rt.models.groot_n17.pipeline_thor.dit_forward` with the AMD
13
+ kernel module passed in — every binding name the bf16 DiT path touches
14
+ (``bf16_nn`` / ``add_bias_bf16`` / ``ada_layer_norm_bf16`` /
15
+ ``layer_norm_no_affine_bf16`` / ``gelu_inplace`` / ``residual_add`` /
16
+ ``concat2_bf16`` / ``relu_inplace_bf16`` / ``silu_inplace_fp16`` /
17
+ ``cast_*`` / ``gpu_copy``) exists in ``flash_rt_amd_kernels``.
18
+
19
+ Reuse strategy: subclass ``_GrootN17FP8BackboneMixin`` +
20
+ ``GrootN17TorchFrontendThor``. The Thor base resolves its kernel module
21
+ and GEMM runners lazily behind ``hasattr(self, "_fvk"/"_gemm")`` guards,
22
+ so this class seeds them with the AMD module at construction and the
23
+ ~1900 validated base lines (weight loading via WEIGHT_SPEC, aux
24
+ contract, calibration bake/save/cache, diffusion modulator precompute,
25
+ graph replay ``infer``) run unmodified. torch on ROCm keeps the "cuda"
26
+ device string, ``torch.cuda.Stream`` and ``torch.cuda.CUDAGraph``
27
+ (hipGraph underneath), so the base's graph machinery is used as-is.
28
+
29
+ AMD-specific overrides:
30
+
31
+ * ``_run_kernel_backbone_fp8`` — AMD pipeline stages + the CDNA4
32
+ attention backend. The llm stage's Q/K/V GEMMs write straight into
33
+ the backend slots (Q 16 heads, K/V the NATIVE 8 KV heads — aiter
34
+ handles GQA internally), so the RTX K/V staging buffers and the
35
+ ``gpu_repeat_interleave_heads`` expand step are dropped.
36
+ * ``_build_dit_attn`` — one :class:`Cdna4GrootN17AttnBackend` serves
37
+ all five sites; the backbone-time instance is reused for the DiT
38
+ when the action-token count matches.
39
+ * ``_setup_cross_kv_kernel`` / ``_cross_kv_fwd`` — the per-frame
40
+ cross-KV refresh writes directly into the backend's per-block K/V
41
+ slots (exact-size prefix views of the padded slots). The gather of
42
+ text/image backbone rows uses ``torch.index_select`` with a
43
+ preallocated ``out=`` (the AMD kernel module has no
44
+ ``embedding_lookup_bf16`` binding); with fixed shapes and no
45
+ allocation it is graph-capture-safe.
46
+ * ``_capture_kernel_dit_graphs`` — same fully-kernelized single-graph
47
+ chain as Thor with two deltas: the Euler update runs as an in-place
48
+ ``torch.add_`` on the persistent action buffer (no
49
+ ``euler_step_bf16_out`` binding on AMD), and — per the aiter
50
+ capture-safety contract — the warmup iterations run ON the capture
51
+ stream (inside its ``torch.cuda.stream`` context) so aiter's
52
+ workspace allocations reach caching-allocator steady state before
53
+ stream capture begins and every torch-side call lands on the same
54
+ stream the raw-int kernels use.
55
+
56
+ The FP8 / FP4 DiT quantization tiers of the base are NOT ported;
57
+ ``_capture_kernel_dit_graphs`` rejects them explicitly.
58
+ """
59
+
60
+ from __future__ import annotations
61
+
62
+ import os
63
+
64
+ import torch
65
+
66
+ from flash_rt.frontends.torch.groot_n17_thor import GrootN17TorchFrontendThor
67
+ from flash_rt.frontends.torch.groot_n17_rtx_fp8 import _GrootN17FP8BackboneMixin
68
+
69
+ _FP16 = torch.float16
70
+ _U8 = torch.uint8
71
+
72
+
73
+ def _fused_epilogue_enabled() -> bool:
74
+ """FVK_AMD_FUSED_EPILOGUE gate (default ON) for BOTH fusions.
75
+
76
+ When on: the FP8 backbone forwards run with ``fused_epilogue=True``
77
+ (bias / bias+GELU in the hipBLASLt FP8 epilogue, host alphas) and the
78
+ DiT runs the AMD-local fused ``dit_forward`` (``bf16_nn_bias`` /
79
+ ``bf16_nn_bias_gelu``). One env flips both for the A/B against the
80
+ decomposed form. Read at setup time (set_prompt / graph build), never
81
+ on the hot path — flipping the env after graphs are built has no
82
+ effect on the built graphs.
83
+ """
84
+ return os.environ.get("FVK_AMD_FUSED_EPILOGUE", "1").strip().lower() \
85
+ not in ("0", "off", "false", "no")
86
+
87
+
88
+ class GrootN17TorchFrontendAmd(_GrootN17FP8BackboneMixin,
89
+ GrootN17TorchFrontendThor):
90
+ """N1.7 CDNA4 frontend: FP8 kernel backbone + bf16 DiT action head."""
91
+
92
+ # The DiT runs unquantized bf16 on AMD. NOTE this is a weaker tier
93
+ # than the Thor/RTX FP8 frontends, which inherit _DIT_USE_FP8 = True
94
+ # and quantize the DiT FFN and fused self-attention QKV as well; that
95
+ # calibration path is not ported to CDNA4 yet.
96
+ _DIT_USE_FP8 = False
97
+
98
+ # Default DiT token count (1 state + 40 action tokens) used to size
99
+ # the shared attention backend at backbone-build time; a differing
100
+ # action_horizon at infer time rebuilds a matching backend.
101
+ _DEFAULT_SA = 41
102
+
103
+ def __init__(
104
+ self,
105
+ checkpoint_path: str,
106
+ *,
107
+ num_views: int = 2,
108
+ embodiment_tag: str = "oxe_droid_relative_eef_relative_joint",
109
+ device: str = "cuda:0",
110
+ ):
111
+ # gfx950-only gate, FIRST: the MFMA tile shapes and FP8 paths in
112
+ # the extension are CDNA4-specific — refuse before touching the
113
+ # checkpoint unless BOTH the visible device arch (e.g.
114
+ # "gfx950:sramecc+:xnack-") and the extension's compile-time
115
+ # gpu_arch are gfx950. Anything else computes garbage, not a
116
+ # fallback (a forced hardware="amd_cdna4" on gfx942 fails here).
117
+ # This runs ahead of super().__init__, which loads weights.
118
+ from flash_rt.amd import flash_rt_amd_kernels as _fvk_gate
119
+ # Compare the base target only ("gfx950" from
120
+ # "gfx950:sramecc+:xnack-"); a prefix test would also accept a
121
+ # future "gfx9500".
122
+ _dev_arch = str(_fvk_gate.device_arch())
123
+ _build_arch = str(dict(_fvk_gate.build_info()).get("gpu_arch",
124
+ "unknown"))
125
+ if not (_dev_arch.split(":", 1)[0] == "gfx950"
126
+ and _build_arch.split(":", 1)[0] == "gfx950"):
127
+ raise RuntimeError(
128
+ "GrootN17TorchFrontendAmd is gfx950-only (CDNA4 / "
129
+ f"MI350-series): running device arch is {_dev_arch!r} and "
130
+ f"the extension was built for gpu_arch {_build_arch!r}. "
131
+ "Rebuild with scripts/amd/build_amd.sh on a gfx950 "
132
+ "machine; other AMD arches are not supported by this "
133
+ "backend.")
134
+
135
+ # The strided-FMHA side-load is a CUDA-only .so (Thor ViT path);
136
+ # the AMD ViT attention runs through the CDNA4 backend instead.
137
+ super().__init__(
138
+ checkpoint_path,
139
+ num_views=num_views,
140
+ embodiment_tag=embodiment_tag,
141
+ device=device,
142
+ load_strided_fmha=False,
143
+ )
144
+ # Seed the kernel-module / GEMM-runner handles with the AMD
145
+ # surface BEFORE any base method runs its lazy
146
+ # ``import flash_rt.flash_rt_kernels`` fallback — every such
147
+ # import in the base is guarded by ``hasattr(self, "_fvk")`` /
148
+ # ``hasattr(self, "_gemm")`` / ``hasattr(self, "_mlp_gemm")``.
149
+ from flash_rt.amd import flash_rt_amd_kernels as fvk
150
+ # (gfx950 gate already ran at the top of __init__.)
151
+ self._fvk = fvk
152
+ self._gemm = fvk.GemmRunner()
153
+ self._mlp_gemm = fvk.GemmRunner()
154
+ # Timed hipBLASLt algorithm selection on the first eager call of
155
+ # each GEMM shape (all first calls happen pre-capture: shadow
156
+ # calibration / set_prompt backbone / graph warmup). gfx950
157
+ # heuristics have known gaps; pi05 measured meaningful wins from
158
+ # timed picks. FLASHRT_FP8_NT_AUTOTUNE=off disables;
159
+ # FLASHRT_FP8_ALGO_POOL sets the candidate pool (default 16 —
160
+ # deeper pools widen run-to-run pick variance).
161
+ import os as _os
162
+ if _os.environ.get("FLASHRT_FP8_NT_AUTOTUNE", "auto").lower() != "off":
163
+ pool = int(_os.environ.get("FLASHRT_FP8_ALGO_POOL", "16"))
164
+ self._gemm.enable_lazy_autotune(pool)
165
+ self._mlp_gemm.enable_lazy_autotune(pool)
166
+
167
+ # ────────────────────────────────────────────────────────────────
168
+ # Attention backend (single instance, all 5 sites)
169
+ # ────────────────────────────────────────────────────────────────
170
+
171
+ def set_hf_processor(self, processor) -> None:
172
+ """Supply the HF processor instead of letting it be loaded.
173
+
174
+ ``denormalize_action`` needs the processor for the
175
+ relative-to-absolute action decode. It is normally built on
176
+ demand from the checkpoint, which reaches the Hub for the
177
+ tokenizer's metadata. Callers that must avoid that lookup —
178
+ offline deployments, or environments hitting the transformers
179
+ 4.57.x bug described in :meth:`_hf_processor` — can build the
180
+ processor themselves and inject it here before inference.
181
+ """
182
+ self._hf_proc_cached = processor
183
+
184
+ def _hf_processor(self):
185
+ """Build the HF processor, with an actionable offline error.
186
+
187
+ Loading the processor resolves its tokenizer by Hub repository
188
+ id, and transformers 4.57.x looks that repository's metadata up
189
+ through ``huggingface_hub.model_info`` without catching
190
+ ``OfflineModeIsEnabled``. With ``HF_HUB_OFFLINE=1`` the load then
191
+ fails even when every file is already cached locally.
192
+
193
+ The library does not work around this: the only interception
194
+ point is ``huggingface_hub``'s module-level function, and
195
+ swapping that out — even temporarily — mutates state shared with
196
+ every other thread in the process. Instead the failure is
197
+ reported with the two remedies that do not.
198
+ """
199
+ if hasattr(self, "_hf_proc_cached"):
200
+ return self._hf_proc_cached
201
+ try:
202
+ return super()._hf_processor()
203
+ except Exception as exc:
204
+ if type(exc).__name__ != "OfflineModeIsEnabled":
205
+ raise
206
+ raise RuntimeError(
207
+ "loading the GROOT processor requires a Hub metadata "
208
+ "lookup that offline mode blocks. This is a transformers "
209
+ "4.57.x issue (the tokenizer loader calls "
210
+ "huggingface_hub.model_info without handling offline "
211
+ "mode), not a missing file. Either allow that one "
212
+ "lookup by clearing HF_HUB_OFFLINE for the load, or "
213
+ "build the processor in your own setup code and pass it "
214
+ "to set_hf_processor() before inference."
215
+ ) from exc
216
+
217
+ def _dit_kv_split(self) -> tuple:
218
+ """(num_text_tokens, num_image_tokens) from the prompt's mask."""
219
+ mask = self._visual_pos_masks
220
+ n_text = int((~mask).sum().item())
221
+ n_image = int(mask.sum().item())
222
+ return n_text, n_image
223
+
224
+ def _build_dit_attn(self, Sa: int) -> None:
225
+ """Bind the CDNA4 backend for the DiT sites.
226
+
227
+ Reuses the backbone-time backend when its DiT token capacity
228
+ matches ``Sa``; otherwise constructs a fresh full backend with
229
+ the same backbone geometry. When the eager torch path has
230
+ already produced exact-size cross K/V tensors
231
+ (``_precompute_dit_cross_kv``), their contents are copied into
232
+ the backend's padded per-block slots; once
233
+ ``_setup_cross_kv_kernel`` rebinds ``_dit_cross_K/V`` to slot
234
+ prefix views, source and destination alias and the copy is
235
+ skipped.
236
+ """
237
+ from flash_rt.amd.hardware.cdna4.attn_backend_groot_n17 import (
238
+ Cdna4GrootN17AttnBackend,
239
+ )
240
+
241
+ Sa = int(Sa)
242
+ n_text, n_image = self._dit_kv_split()
243
+ kv_max = max(n_text, n_image)
244
+
245
+ attn = getattr(self, "_kbb_attn", None)
246
+ if attn is None or int(getattr(self, "_kbb_attn_sa", -1)) != Sa:
247
+ attn = Cdna4GrootN17AttnBackend(
248
+ num_vit_views=int(getattr(self, "_num_vit_views",
249
+ self.num_views)),
250
+ vit_seq=int(self._S_vit),
251
+ llm_seq=int(self.Se),
252
+ vl_self_attn_seq=int(self.Se),
253
+ sa=Sa,
254
+ dit_kv_seq=kv_max,
255
+ device=self.device,
256
+ )
257
+ if hasattr(self, "_dit_cross_K"):
258
+ for j, (k_src, v_src) in enumerate(
259
+ zip(self._dit_cross_K, self._dit_cross_V)):
260
+ rows = int(k_src.shape[0])
261
+ k_dst = attn.dit_cross_K[j].view(kv_max, -1)[:rows]
262
+ v_dst = attn.dit_cross_V[j].view(kv_max, -1)[:rows]
263
+ if k_dst.data_ptr() != k_src.data_ptr():
264
+ k_dst.copy_(k_src)
265
+ if v_dst.data_ptr() != v_src.data_ptr():
266
+ v_dst.copy_(v_src)
267
+ self._dit_attn = attn
268
+
269
+ # ────────────────────────────────────────────────────────────────
270
+ # Kernelized DiT cross-KV (writes straight into backend slots)
271
+ # ────────────────────────────────────────────────────────────────
272
+
273
+ def _setup_cross_kv_kernel(self) -> None:
274
+ """Persistent cross-KV buffers over the CDNA4 backend slots.
275
+
276
+ Mirrors the Thor base with one structural delta: the per-block
277
+ K/V destinations are exact-size prefix views of the backend's
278
+ padded ``dit_cross_K/V[j]`` slots (row-major ``(kv, NH*HD)``
279
+ prefix of ``(kv_max, NH, HD)``), so the projection GEMMs land
280
+ their output exactly where ``attn.run("dit_cross", ...)`` reads.
281
+ """
282
+ if hasattr(self, "_ck_text_idx"):
283
+ return
284
+ if not hasattr(self, "_dit_attn"):
285
+ raise RuntimeError(
286
+ "_setup_cross_kv_kernel requires the DiT attention backend; "
287
+ "call _build_dit_attn first")
288
+ dev = self.device
289
+ S = self.Se
290
+ mask = self._visual_pos_masks
291
+ self._ck_text_idx = torch.where(~mask)[0].to(torch.int64).contiguous()
292
+ self._ck_image_idx = torch.where(mask)[0].to(torch.int64).contiguous()
293
+ nt = int(self._ck_text_idx.numel())
294
+ ni = int(self._ck_image_idx.numel())
295
+ self._ck_nt, self._ck_ni = nt, ni
296
+ bf = torch.bfloat16
297
+ # Per-frame backbone input (fp16, copied in before replay) + bf16 cast
298
+ self._ck_bb_src = torch.empty(S, 2048, dtype=torch.float16, device=dev)
299
+ self._ck_bb = torch.empty(S, 2048, dtype=bf, device=dev)
300
+ self._ck_text_src = torch.empty(nt, 2048, dtype=bf, device=dev)
301
+ self._ck_image_src = torch.empty(ni, 2048, dtype=bf, device=dev)
302
+
303
+ # Cross K/V destinations = exact-size prefix views of the backend
304
+ # slots. Block j maps to full-layer index li = 2j; text-target
305
+ # blocks (li % 4 == 0) are the even j.
306
+ attn = self._dit_attn
307
+ kv_max = int(attn.dit_cross_K[0].shape[0])
308
+ D = 1536
309
+ self._dit_cross_K = [
310
+ attn.dit_cross_K[j].view(kv_max, D)[: (nt if j % 2 == 0 else ni)]
311
+ for j in range(16)]
312
+ self._dit_cross_V = [
313
+ attn.dit_cross_V[j].view(kv_max, D)[: (nt if j % 2 == 0 else ni)]
314
+ for j in range(16)]
315
+
316
+ # Seed from the current backbone (eager) so a non-graph consumer
317
+ # sees valid K/V immediately.
318
+ self._ck_bb_src.copy_(self._backbone_features.reshape(S, 2048).half())
319
+ self._cross_kv_fwd(0)
320
+
321
+ def _cross_kv_fwd(self, s: int) -> None:
322
+ """Cross-KV forward over the persistent buffers (graph-safe).
323
+
324
+ Reads ``_ck_bb_src`` (current backbone, fp16), writes the
325
+ backend-slot K/V prefixes. The text/image row gather runs as
326
+ ``torch.index_select`` with preallocated ``out=`` buffers — the
327
+ AMD kernel module has no ``embedding_lookup_bf16`` binding. The
328
+ torch calls land on torch's current stream; callers keep it
329
+ consistent with the raw ``s`` int (stream 0 eagerly, or the
330
+ capture stream via its ``torch.cuda.stream`` context).
331
+ """
332
+ K = self._fvk
333
+ mg = self._mlp_gemm
334
+ S = self.Se
335
+ nt, ni = self._ck_nt, self._ck_ni
336
+ fused = _fused_epilogue_enabled()
337
+ K.cast_fp16_to_bf16(
338
+ self._ck_bb_src.data_ptr(), self._ck_bb.data_ptr(),
339
+ S * 2048, int(s))
340
+ torch.index_select(
341
+ self._ck_bb, 0, self._ck_text_idx, out=self._ck_text_src)
342
+ torch.index_select(
343
+ self._ck_bb, 0, self._ck_image_idx, out=self._ck_image_src)
344
+ for j in range(16):
345
+ li = 2 * j
346
+ text = (li % 4 == 0)
347
+ N = nt if text else ni
348
+ src = self._ck_text_src if text else self._ck_image_src
349
+ k_dst = self._dit_cross_K[j]
350
+ v_dst = self._dit_cross_V[j]
351
+ k_w = self._dit_k_w[li]
352
+ k_b = self._dit_k_b[li]
353
+ v_w = self._dit_v_w[li]
354
+ v_b = self._dit_v_b[li]
355
+ if fused:
356
+ # AMD FUSED: bf16_nn + add_bias_bf16 → bf16_nn_bias (cross K/V).
357
+ mg.bf16_nn_bias(src.data_ptr(), k_w.data_ptr(),
358
+ k_dst.data_ptr(), k_b.data_ptr(),
359
+ N, 1536, 2048, int(s))
360
+ mg.bf16_nn_bias(src.data_ptr(), v_w.data_ptr(),
361
+ v_dst.data_ptr(), v_b.data_ptr(),
362
+ N, 1536, 2048, int(s))
363
+ else:
364
+ mg.bf16_nn(src.data_ptr(), k_w.data_ptr(), k_dst.data_ptr(),
365
+ N, 1536, 2048, int(s))
366
+ K.add_bias_bf16(k_dst.data_ptr(), k_b.data_ptr(), N, 1536,
367
+ int(s))
368
+ mg.bf16_nn(src.data_ptr(), v_w.data_ptr(), v_dst.data_ptr(),
369
+ N, 1536, 2048, int(s))
370
+ K.add_bias_bf16(v_dst.data_ptr(), v_b.data_ptr(), N, 1536,
371
+ int(s))
372
+
373
+ # ────────────────────────────────────────────────────────────────
374
+ # Fully-kernelized DiT graph (bf16, single combined graph)
375
+ # ────────────────────────────────────────────────────────────────
376
+
377
+ def _pack_smallm_dit_weights(self) -> None:
378
+ """MFMA-pack the DiT (41, 1536, 1536) projection weights ONCE at
379
+ graph-build time (never per frame) for the gfx950 small-M packed
380
+ bf16 kernel (csrc/amd/gemm/smallm_mfma_bf16.h).
381
+
382
+ Measured routing (standalone gate vs autotuned hipBLASLt): the
383
+ packed kernel wins ONLY on the square D→D shape — bias 1.41x,
384
+ bias_res 1.12x — and loses on ff1/ff2, so only the projection
385
+ weights feeding those sites get packed copies: Q (all 32
386
+ layers), K/V (the 16 self blocks — the cross blocks' K/V
387
+ weights are the (2048, 1536) cross-KV projections, a different
388
+ shape consumed by ``_setup_cross_kv_kernel``), O (all 32).
389
+
390
+ Packed copies land in ``self._dit_smallm_packed`` keyed
391
+ ``{q,k,v,o}_w_{li}`` (pointer ints, matching the pipeline
392
+ ``weights`` lists); the tensors are kept alive in
393
+ ``self._dit_smallm_store``. Originals are KEPT — they serve the
394
+ FVK_AMD_DIT_GEMM=hipblaslt fallback and the cross-KV consumer —
395
+ at ~432 MB extra VRAM for the 96 packed copies.
396
+ FVK_AMD_DIT_GEMM=hipblaslt (read once here, setup time) skips
397
+ the packing entirely; the AMD ``dit_forward`` reads the same
398
+ env for routing.
399
+ """
400
+ if getattr(self, "_dit_smallm_packed", None) is not None:
401
+ return
402
+ self._dit_smallm_packed: dict = {}
403
+ self._dit_smallm_store: list = []
404
+ # FVK_AMD_DIT_GEMM: "smallm" (default) = pack + route the D→D
405
+ # projections to the MFMA packed kernel; "hipblaslt" = library path.
406
+ if os.environ.get("FVK_AMD_DIT_GEMM", "smallm").strip().lower() \
407
+ != "smallm":
408
+ return
409
+ with torch.no_grad():
410
+ for li in range(32):
411
+ sites = [("q_w", self._dit_q_w), ("o_w", self._dit_o_w)]
412
+ if li % 2 == 1: # self blocks only (K/V)
413
+ sites += [("k_w", self._dit_k_w),
414
+ ("v_w", self._dit_v_w)]
415
+ for key, wl in sites:
416
+ W = wl[li] # (K, N) row-major bf16, both 1536
417
+ Kd, Nd = W.shape
418
+ # Per-lane consumption order — the documented
419
+ # one-liner from smallm_mfma_bf16.h.
420
+ wp = (W.view(Kd // 32, 4, 8, Nd // 16, 16)
421
+ .permute(3, 0, 1, 4, 2).contiguous())
422
+ self._dit_smallm_store.append(wp)
423
+ self._dit_smallm_packed[f"{key}_{li}"] = wp.data_ptr()
424
+
425
+ def _capture_kernel_dit_graphs(self, num_inference_timesteps: int = 4,
426
+ action_horizon: int = 40) -> None:
427
+ """Capture the per-frame action-head chain as ONE HIP graph.
428
+
429
+ AMD adaptation of the Thor base method: same buffer set, same
430
+ per-step closures, same combined cross-KV + state-encode +
431
+ 4-step-DiT graph, with three deltas (see module docstring):
432
+ bf16-only DiT (no FP8/FP4 splice), a torch in-place Euler add,
433
+ and warmup ON the capture stream per the aiter capture-safety
434
+ contract.
435
+ """
436
+ from flash_rt.models.groot_n17 import pipeline_thor
437
+
438
+ # FVK_AMD_FUSED_EPILOGUE flips BOTH fusions (backbone + DiT) for
439
+ # the A/B: on → the AMD-local dit_forward (bf16_nn_bias /
440
+ # bf16_nn_bias_gelu fused epilogues); off → the byte-identical
441
+ # decomposed pipeline_thor.dit_forward. Only dit_forward is
442
+ # AMD-local; embodiment_* stages keep coming from pipeline_thor.
443
+ fused_ep = _fused_epilogue_enabled()
444
+ if fused_ep:
445
+ from flash_rt.amd.models.groot_n17 import pipeline as _amd_pipeline
446
+ dit_forward = _amd_pipeline.dit_forward
447
+ else:
448
+ dit_forward = pipeline_thor.dit_forward
449
+
450
+ if getattr(self, "_DIT_USE_FP8", False) or \
451
+ getattr(self, "_DIT_QUANT", "fp8") == "fp4":
452
+ raise NotImplementedError(
453
+ "the AMD CDNA4 N1.7 frontend runs the DiT bf16; the FP8/FP4 "
454
+ "DiT quantization tiers are not ported")
455
+
456
+ Sa = action_horizon + 1
457
+ if not hasattr(self, "_dit_attn"):
458
+ self._build_dit_attn(Sa)
459
+ self._setup_cross_kv_kernel()
460
+ if not hasattr(self, "_infer_bufs"):
461
+ self._allocate_infer_buffers(action_horizon)
462
+ self._prepare_kernel_dit(num_inference_timesteps)
463
+ self._allocate_kernel_dit_buffers(action_horizon)
464
+ if fused_ep:
465
+ # Setup-time MFMA packing for the smallm D→D projection
466
+ # routing in the AMD dit_forward (FVK_AMD_DIT_GEMM).
467
+ self._pack_smallm_dit_weights()
468
+
469
+ K = self._fvk
470
+ mg = self._mlp_gemm
471
+ w = self._kw
472
+ bufs = self._infer_bufs
473
+ dit_h = bufs["dit_h"].data_ptr()
474
+ Skv_text = int(self._dit_cross_K[0].shape[0])
475
+ Skv_image = int(self._dit_cross_K[1].shape[0])
476
+ dims = {"Sa": Sa, "D": 1536, "FF": 6144,
477
+ "Skv_text": Skv_text, "Skv_image": Skv_image}
478
+ bp = {"h": dit_h, "xn": bufs["dit_xn"].data_ptr(),
479
+ "o_proj_out": bufs["dit_o_proj_out"].data_ptr(),
480
+ "ff_proj_out": bufs["dit_ff_proj_out"].data_ptr()}
481
+ dt = 1.0 / num_inference_timesteps
482
+ # decode reads the action rows (1..Sa) of the (Sa, 1024) output_proj
483
+ hout_dec = self._k_hout.data_ptr() + 1024 * 2 # skip state row
484
+
485
+ def _dit_weights(step):
486
+ d = {"scale_msa": [t.data_ptr() for t in self._step_scales[step]],
487
+ "shift_msa": [t.data_ptr() for t in self._step_shifts[step]]}
488
+ for key, attr in (("q_w", "_dit_q_w"), ("q_b", "_dit_q_b"),
489
+ ("k_w", "_dit_k_w"), ("k_b", "_dit_k_b"),
490
+ ("v_w", "_dit_v_w"), ("v_b", "_dit_v_b"),
491
+ ("o_w", "_dit_o_w"), ("o_b", "_dit_o_b"),
492
+ ("ff_proj_w", "_dit_ff_proj_w"),
493
+ ("ff_proj_b", "_dit_ff_proj_b"),
494
+ ("ff_down_w", "_dit_ff_down_w"),
495
+ ("ff_down_b", "_dit_ff_down_b")):
496
+ d[key] = [t.data_ptr() for t in getattr(self, attr)]
497
+ if fused_ep:
498
+ # MFMA-packed q/k/v/o copies for the smallm routing in
499
+ # the AMD dit_forward (same dict for every step —
500
+ # weights are step-invariant). Empty dict when
501
+ # FVK_AMD_DIT_GEMM=hipblaslt skipped the packing.
502
+ d["smallm_packed"] = self._dit_smallm_packed
503
+ return d
504
+
505
+ step_weights = [_dit_weights(s) for s in range(num_inference_timesteps)]
506
+
507
+ def _state_fwd(s):
508
+ if fused_ep:
509
+ # AMD FUSED: bf16_nn + add_bias_bf16 → bf16_nn_bias (state l1).
510
+ mg.bf16_nn_bias(self._k_state_in.data_ptr(),
511
+ w["st_l1"].data_ptr(),
512
+ self._k_st_h1.data_ptr(),
513
+ w["st_l1b"].data_ptr(), 1, 1024, 132, s)
514
+ else:
515
+ mg.bf16_nn(self._k_state_in.data_ptr(), w["st_l1"].data_ptr(),
516
+ self._k_st_h1.data_ptr(), 1, 1024, 132, s)
517
+ K.add_bias_bf16(self._k_st_h1.data_ptr(),
518
+ w["st_l1b"].data_ptr(), 1, 1024, s)
519
+ K.relu_inplace_bf16(self._k_st_h1.data_ptr(), 1024, s)
520
+ if fused_ep:
521
+ # AMD FUSED: bf16_nn + add_bias_bf16 → bf16_nn_bias (state l2).
522
+ mg.bf16_nn_bias(self._k_st_h1.data_ptr(),
523
+ w["st_l2"].data_ptr(),
524
+ self._k_state_feat.data_ptr(),
525
+ w["st_l2b"].data_ptr(), 1, 1536, 1024, s)
526
+ else:
527
+ mg.bf16_nn(self._k_st_h1.data_ptr(), w["st_l2"].data_ptr(),
528
+ self._k_state_feat.data_ptr(), 1, 1536, 1024, s)
529
+ K.add_bias_bf16(self._k_state_feat.data_ptr(),
530
+ w["st_l2b"].data_ptr(), 1, 1536, s)
531
+
532
+ def _ae_fwd(step, s):
533
+ # action_encode: W1 (no act) → cat[a_emb, tau] → W2 → SiLU → W3,
534
+ # add pos, then fill dit_h ([0]=state, [1:]=action features).
535
+ T = action_horizon
536
+ if fused_ep:
537
+ # AMD FUSED: bf16_nn + add_bias_bf16 → bf16_nn_bias (ae W1).
538
+ mg.bf16_nn_bias(self._k_actions.data_ptr(),
539
+ w["ae_W1"].data_ptr(),
540
+ self._k_ae_aemb.data_ptr(),
541
+ w["ae_b1"].data_ptr(), T, 1536, 132, s)
542
+ else:
543
+ mg.bf16_nn(self._k_actions.data_ptr(), w["ae_W1"].data_ptr(),
544
+ self._k_ae_aemb.data_ptr(), T, 1536, 132, s)
545
+ K.add_bias_bf16(self._k_ae_aemb.data_ptr(),
546
+ w["ae_b1"].data_ptr(), T, 1536, s)
547
+ K.concat2_bf16(self._k_ae_aemb.data_ptr(),
548
+ self._k_tau[step].data_ptr(),
549
+ self._k_ae_concat.data_ptr(), T, 1536, 1536, s)
550
+ if fused_ep:
551
+ # AMD FUSED: bf16_nn + add_bias_bf16 → bf16_nn_bias (ae W2).
552
+ mg.bf16_nn_bias(self._k_ae_concat.data_ptr(),
553
+ w["ae_W2"].data_ptr(),
554
+ self._k_ae_i2.data_ptr(),
555
+ w["ae_b2"].data_ptr(), T, 1536, 3072, s)
556
+ else:
557
+ mg.bf16_nn(self._k_ae_concat.data_ptr(), w["ae_W2"].data_ptr(),
558
+ self._k_ae_i2.data_ptr(), T, 1536, 3072, s)
559
+ K.add_bias_bf16(self._k_ae_i2.data_ptr(),
560
+ w["ae_b2"].data_ptr(), T, 1536, s)
561
+ K.cast_bf16_to_fp16(self._k_ae_i2.data_ptr(),
562
+ self._k_ae_i2f.data_ptr(), T * 1536, s)
563
+ K.silu_inplace_fp16(self._k_ae_i2f.data_ptr(), T * 1536, s)
564
+ K.cast_fp16_to_bf16(self._k_ae_i2f.data_ptr(),
565
+ self._k_ae_i2.data_ptr(), T * 1536, s)
566
+ if fused_ep:
567
+ # AMD FUSED: bf16_nn + add_bias_bf16 → bf16_nn_bias (ae W3).
568
+ mg.bf16_nn_bias(self._k_ae_i2.data_ptr(),
569
+ w["ae_W3"].data_ptr(),
570
+ self._k_ae_out.data_ptr(),
571
+ w["ae_b3"].data_ptr(), T, 1536, 1536, s)
572
+ else:
573
+ mg.bf16_nn(self._k_ae_i2.data_ptr(), w["ae_W3"].data_ptr(),
574
+ self._k_ae_out.data_ptr(), T, 1536, 1536, s)
575
+ K.add_bias_bf16(self._k_ae_out.data_ptr(),
576
+ w["ae_b3"].data_ptr(), T, 1536, s)
577
+ K.residual_add(self._k_ae_out.data_ptr(), self._k_pos.data_ptr(),
578
+ T * 1536, s)
579
+ K.gpu_copy(dit_h, self._k_state_feat.data_ptr(), 1536 * 2, s)
580
+ K.gpu_copy(dit_h + 1536 * 2, self._k_ae_out.data_ptr(),
581
+ T * 1536 * 2, s)
582
+
583
+ def _post_fwd(step, s):
584
+ # output projection (AdaLN → proj_out_2) + action_decode + Euler.
585
+ T = action_horizon
586
+ K.ada_layer_norm_bf16(dit_h, self._k_oproj_scale[step].data_ptr(),
587
+ self._k_oproj_shift[step].data_ptr(),
588
+ self._k_hmod.data_ptr(), Sa, 1536, 1e-5, s)
589
+ if fused_ep:
590
+ # AMD FUSED: bf16_nn + add_bias_bf16 → bf16_nn_bias (po2).
591
+ mg.bf16_nn_bias(self._k_hmod.data_ptr(), w["po2"].data_ptr(),
592
+ self._k_hout.data_ptr(),
593
+ w["po2b"].data_ptr(), Sa, 1024, 1536, s)
594
+ # AMD FUSED: bf16_nn + add_bias_bf16 → bf16_nn_bias (dec l1).
595
+ mg.bf16_nn_bias(hout_dec, w["dec_l1"].data_ptr(),
596
+ self._k_dec_h.data_ptr(),
597
+ w["dec_l1b"].data_ptr(), T, 1024, 1024, s)
598
+ else:
599
+ mg.bf16_nn(self._k_hmod.data_ptr(), w["po2"].data_ptr(),
600
+ self._k_hout.data_ptr(), Sa, 1024, 1536, s)
601
+ K.add_bias_bf16(self._k_hout.data_ptr(), w["po2b"].data_ptr(),
602
+ Sa, 1024, s)
603
+ mg.bf16_nn(hout_dec, w["dec_l1"].data_ptr(),
604
+ self._k_dec_h.data_ptr(), T, 1024, 1024, s)
605
+ K.add_bias_bf16(self._k_dec_h.data_ptr(),
606
+ w["dec_l1b"].data_ptr(), T, 1024, s)
607
+ K.relu_inplace_bf16(self._k_dec_h.data_ptr(), T * 1024, s)
608
+ if fused_ep:
609
+ # AMD FUSED: bf16_nn + add_bias_bf16 → bf16_nn_bias (dec l2).
610
+ mg.bf16_nn_bias(self._k_dec_h.data_ptr(),
611
+ w["dec_l2"].data_ptr(),
612
+ self._k_vel.data_ptr(),
613
+ w["dec_l2b"].data_ptr(), T, 132, 1024, s)
614
+ else:
615
+ mg.bf16_nn(self._k_dec_h.data_ptr(), w["dec_l2"].data_ptr(),
616
+ self._k_vel.data_ptr(), T, 132, 1024, s)
617
+ K.add_bias_bf16(self._k_vel.data_ptr(),
618
+ w["dec_l2b"].data_ptr(), T, 132, s)
619
+ # Euler update: actions += dt * velocity, in place on the
620
+ # persistent buffers. The AMD module has no euler_step
621
+ # binding; an in-place torch add is capture-safe (fixed
622
+ # shapes, no allocation) and lands on the capture stream via
623
+ # the surrounding torch.cuda.stream context.
624
+ self._k_actions.add_(self._k_vel, alpha=dt)
625
+
626
+ def _step_fwd(step, s):
627
+ _ae_fwd(step, s)
628
+ dit_forward(
629
+ gemm=self._gemm, fvk=K, bufs=bp, weights=step_weights[step],
630
+ dims=dims, attn=self._dit_attn, stream=s)
631
+ _post_fwd(step, s)
632
+
633
+ self._kdit_fwd = (_state_fwd, _step_fwd)
634
+ self._k_nsteps = num_inference_timesteps
635
+
636
+ def _dit_all(s):
637
+ self._cross_kv_fwd(s)
638
+ _state_fwd(s)
639
+ for step in range(num_inference_timesteps):
640
+ _step_fwd(step, s)
641
+
642
+ # aiter capture-safety contract (see the CDNA4 backend docstring):
643
+ # aiter may allocate LSE/workspace through torch's caching
644
+ # allocator per call, so the warmup iterations must run ON the
645
+ # capture stream (allocator steady state per stream) and every
646
+ # torch-side call must land on that same stream. The warmup also
647
+ # primes the hipBLASLt GemmRunner algo caches for every captured
648
+ # GEMM shape.
649
+ stream = torch.cuda.Stream()
650
+ stream.wait_stream(torch.cuda.current_stream())
651
+ with torch.cuda.stream(stream):
652
+ s_int = stream.cuda_stream
653
+ for _ in range(3):
654
+ _dit_all(s_int)
655
+ torch.cuda.synchronize()
656
+
657
+ graph = torch.cuda.CUDAGraph()
658
+ with torch.cuda.stream(stream):
659
+ graph.capture_begin()
660
+ _dit_all(stream.cuda_stream)
661
+ graph.capture_end()
662
+ torch.cuda.current_stream().wait_stream(stream)
663
+ torch.cuda.synchronize()
664
+ self._k_dit_graph = graph
665
+
666
+ # ────────────────────────────────────────────────────────────────
667
+ # FP8 kernel backbone (AMD pipeline, native-GQA llm slots)
668
+ # ────────────────────────────────────────────────────────────────
669
+
670
+ def _run_kernel_backbone_fp8(self, aux: dict) -> "torch.Tensor":
671
+ """ViT → DeepStack → LLM → vlln → VL-self-attn on AMD FP8 kernels.
672
+
673
+ Mirror of the RTX mixin method retargeted at
674
+ :mod:`flash_rt.amd.models.groot_n17.pipeline` and the CDNA4
675
+ attention backend. The llm stage's Q/K/V descale GEMMs write
676
+ straight into the backend slots (Q 16 heads, K/V the native
677
+ 8 KV heads); no K/V staging buffers and no K_exp/V_exp scratch
678
+ are allocated — aiter consumes the 8 KV heads natively.
679
+ """
680
+ from flash_rt.amd.models.groot_n17 import pipeline as P
681
+ from flash_rt.amd.hardware.cdna4.attn_backend_groot_n17 import (
682
+ Cdna4GrootN17AttnBackend,
683
+ )
684
+
685
+ fvkm, gemm = self._fvk, self._gemm
686
+ dev = self.device
687
+ Sv, nv, Se = self._S_vit, self._num_vit_views, self.Se
688
+
689
+ keep: list = []
690
+ self._kbb_keep = keep
691
+
692
+ def K(t):
693
+ keep.append(t)
694
+ return t
695
+
696
+ def buf(*shape):
697
+ return K(torch.empty(*shape, dtype=_FP16, device=dev))
698
+
699
+ def buf8(*shape):
700
+ return K(torch.empty(*shape, dtype=_U8, device=dev))
701
+
702
+ def wsc(val):
703
+ """Upload a host weight scale to a device fp32 scalar; keep ref."""
704
+ t = K(torch.tensor([float(val)], dtype=torch.float32, device=dev))
705
+ return t.data_ptr()
706
+
707
+ def adv(dev_list):
708
+ """Device act-scale scalar tensors → list of int ptrs."""
709
+ return [t.data_ptr() for t in dev_list]
710
+
711
+ # ── FUSED-EPILOGUE tier (FVK_AMD_FUSED_EPILOGUE, default on) ──
712
+ # Host alphas for the fused fp8_nn_bias / fp8_nn_gelu_bias calls.
713
+ # _bake_calibration (already run via _ensure_act_scales) composes
714
+ # them as python floats: alpha = act_scale × w_scale per site per
715
+ # layer. Key names parallel the scales_dev dicts. The llm stage is
716
+ # biasless (Qwen3) and always stays on descale GEMMs; its
717
+ # fused_epilogue flag fuses the norm/residual+quantize chains
718
+ # instead (rms_norm_fp8_fp16 / residual_add_rms_norm_fp8_fp16 —
719
+ # no alphas needed).
720
+ fused = _fused_epilogue_enabled()
721
+ if fused:
722
+ vit_alphas = {
723
+ "act_qkv": [float(a) for a in self._vit_alpha_q],
724
+ "act_o": [float(a) for a in self._vit_alpha_o],
725
+ "act_fc1": [float(a) for a in self._vit_alpha_fc1],
726
+ "act_fc2": [float(a) for a in self._vit_alpha_fc2],
727
+ }
728
+ ds_alphas = {
729
+ "act_fc1": [float(a) for a in self._dsm_alpha_fc1],
730
+ "act_fc2": [float(a) for a in self._dsm_alpha_fc2],
731
+ }
732
+ # vlsa Q/K/V have separate weight scales → per-layer 3-tuples.
733
+ vlsa_alphas = {
734
+ "act_qkv": [
735
+ (float(q), float(k), float(v))
736
+ for q, k, v in zip(self._vlsa_alpha_q,
737
+ self._vlsa_alpha_k,
738
+ self._vlsa_alpha_v)],
739
+ "act_o": [float(a) for a in self._vlsa_alpha_o],
740
+ "act_fc1": [float(a) for a in self._vlsa_alpha_fc1],
741
+ "act_fc2": [float(a) for a in self._vlsa_alpha_fc2],
742
+ }
743
+ else:
744
+ vit_alphas = ds_alphas = vlsa_alphas = None
745
+
746
+ # One backend for the whole model: backbone sites now, DiT sites
747
+ # at first infer (reused by _build_dit_attn when Sa matches).
748
+ n_text, n_image = self._dit_kv_split()
749
+ sa = int(self._DEFAULT_SA)
750
+ attn = Cdna4GrootN17AttnBackend(
751
+ num_vit_views=nv, vit_seq=Sv, llm_seq=Se, vl_self_attn_seq=Se,
752
+ sa=sa, dit_kv_seq=max(n_text, n_image), device=dev)
753
+ self._kbb_attn = attn
754
+ self._kbb_attn_sa = sa
755
+
756
+ # ═══ ViT (24L) ═══
757
+ vit_h = buf(Sv, 1024)
758
+ vit_h.copy_(aux["pixel_features"].to(dev).half().reshape(Sv, 1024))
759
+ vit_bufs = {"h": vit_h.data_ptr(), "xn": buf(Sv, 1024).data_ptr(),
760
+ "xn_fp8": buf8(Sv, 1024).data_ptr(),
761
+ "o_proj_out": buf(Sv, 1024).data_ptr(),
762
+ "fc1_out": buf(Sv, 4096).data_ptr(),
763
+ "fc1_fp8": buf8(Sv, 4096).data_ptr()}
764
+ vw = {k: [] for k in (
765
+ "norm1_w", "norm1_b", "norm2_w", "norm2_b", "q_w", "q_b",
766
+ "k_w", "k_b", "v_w", "v_b", "o_w", "o_b", "fc1_w", "fc1_b",
767
+ "fc2_w", "fc2_b", "q_ws", "k_ws", "v_ws", "o_ws",
768
+ "fc1_ws", "fc2_ws")}
769
+ vw["cos"] = self._vit_cos.data_ptr()
770
+ vw["sin"] = self._vit_sin.data_ptr()
771
+ for li in range(24):
772
+ qkv = self._vit_qkv_w[li] # fp8 (1024, 3072) [K, 3N]
773
+ b = self._vit_qkv_b[li] # (3072,)
774
+ q = K(qkv[:, :1024].contiguous())
775
+ kk = K(qkv[:, 1024:2048].contiguous())
776
+ v = K(qkv[:, 2048:].contiguous())
777
+ qb = K(b[:1024].contiguous())
778
+ kb = K(b[1024:2048].contiguous())
779
+ vb = K(b[2048:].contiguous())
780
+ qkv_ws = wsc(self._vit_alpha[li * 4 + 0])
781
+ vw["norm1_w"].append(self._vit_ln1_w[li].data_ptr())
782
+ vw["norm1_b"].append(self._vit_ln1_b[li].data_ptr())
783
+ vw["norm2_w"].append(self._vit_ln2_w[li].data_ptr())
784
+ vw["norm2_b"].append(self._vit_ln2_b[li].data_ptr())
785
+ vw["q_w"].append(q.data_ptr()); vw["q_b"].append(qb.data_ptr())
786
+ vw["k_w"].append(kk.data_ptr()); vw["k_b"].append(kb.data_ptr())
787
+ vw["v_w"].append(v.data_ptr()); vw["v_b"].append(vb.data_ptr())
788
+ vw["q_ws"].append(qkv_ws)
789
+ vw["k_ws"].append(qkv_ws)
790
+ vw["v_ws"].append(qkv_ws)
791
+ vw["o_w"].append(self._vit_o_w[li].data_ptr())
792
+ vw["o_b"].append(self._vit_o_b[li].data_ptr())
793
+ vw["o_ws"].append(wsc(self._vit_alpha[li * 4 + 1]))
794
+ vw["fc1_w"].append(self._vit_fc1_w[li].data_ptr())
795
+ vw["fc1_b"].append(self._vit_fc1_b[li].data_ptr())
796
+ vw["fc1_ws"].append(wsc(self._vit_alpha[li * 4 + 2]))
797
+ vw["fc2_w"].append(self._vit_fc2_w[li].data_ptr())
798
+ vw["fc2_b"].append(self._vit_fc2_b[li].data_ptr())
799
+ vw["fc2_ws"].append(wsc(self._vit_alpha[li * 4 + 3]))
800
+ vit_scales = {
801
+ "act_qkv": adv(self._vit_act_qkv_dev),
802
+ "act_o": adv(self._vit_act_o_dev),
803
+ "act_fc1": adv(self._vit_act_fc1_dev),
804
+ "act_fc2": adv(self._vit_act_fc2_dev)}
805
+
806
+ tap_layers = (5, 11, 17)
807
+ tap_bufs = {l: buf(Sv, 1024) for l in tap_layers}
808
+ scell = [0]
809
+ self._kbb_scell = scell
810
+
811
+ def mk_cb(l):
812
+ def cb(h_ptr):
813
+ fvkm.gpu_copy(
814
+ tap_bufs[l].data_ptr(), int(h_ptr), Sv * 1024 * 2,
815
+ scell[0])
816
+ return cb
817
+ dcap = [mk_cb(l) for l in tap_layers]
818
+
819
+ P.qwen3vl_vit_forward(
820
+ gemm=gemm, fvk=fvkm, bufs=vit_bufs, weights=vw,
821
+ scales_dev=vit_scales,
822
+ dims={"S": Sv, "D": 1024, "NH": 16, "HD": 64,
823
+ "ff_inner": 4096, "Sper_view": Sv // nv},
824
+ attn=attn, deepstack_taps=tap_layers, deepstack_capture=dcap,
825
+ fused_epilogue=fused, alphas=vit_alphas)
826
+
827
+ # ═══ DeepStack (3 mergers) ═══
828
+ Nout = Sv // 4
829
+ ds_out = [buf(Nout, 2048) for _ in range(3)]
830
+ dsw = {k: [] for k in ("norm_w", "norm_b", "fc1_w", "fc1_b",
831
+ "fc2_w", "fc2_b", "fc1_ws", "fc2_ws")}
832
+ for j in range(3):
833
+ dsw["norm_w"].append(getattr(self, f"_dsm{j}_norm_w").data_ptr())
834
+ dsw["norm_b"].append(getattr(self, f"_dsm{j}_norm_b").data_ptr())
835
+ dsw["fc1_w"].append(getattr(self, f"_dsm{j}_fc1_w").data_ptr())
836
+ dsw["fc1_b"].append(getattr(self, f"_dsm{j}_fc1_b").data_ptr())
837
+ dsw["fc1_ws"].append(wsc(self._dsm_alpha[j * 2 + 0]))
838
+ dsw["fc2_w"].append(getattr(self, f"_dsm{j}_fc2_w").data_ptr())
839
+ dsw["fc2_b"].append(getattr(self, f"_dsm{j}_fc2_b").data_ptr())
840
+ dsw["fc2_ws"].append(wsc(self._dsm_alpha[j * 2 + 1]))
841
+ ds_scales = {"act_fc1": adv(self._dsm_act_fc1_dev),
842
+ "act_fc2": adv(self._dsm_act_fc2_dev)}
843
+ ds_bufs = {"in": [tap_bufs[l].data_ptr() for l in tap_layers],
844
+ "ln_out": buf(Nout, 4096).data_ptr(),
845
+ "fp8_scratch": buf8(Nout, 4096).data_ptr(),
846
+ "fc1_out": buf(Nout, 4096).data_ptr(),
847
+ "out": [t.data_ptr() for t in ds_out]}
848
+ ds_dims = {"Nin": Sv, "Din": 1024, "Nout": Nout,
849
+ "Dmid": 4096, "Dout": 2048}
850
+ P.deepstack_merge_forward(
851
+ gemm=gemm, fvk=fvkm, bufs=ds_bufs,
852
+ weights=dsw, scales_dev=ds_scales, dims=ds_dims,
853
+ fused_epilogue=fused, alphas=ds_alphas)
854
+
855
+ # DeepStack inject buffers (S, D) — zero except visual positions.
856
+ mask = self._visual_pos_masks
857
+ vis_idx = K(mask.reshape(-1).nonzero(as_tuple=True)[0].to(torch.long))
858
+ inject = [0] * 16
859
+ injb = []
860
+ for j in range(3):
861
+ ib = buf(Se, 2048)
862
+ ib.zero_()
863
+ ib.index_copy_(0, vis_idx, ds_out[j])
864
+ inject[j] = ib.data_ptr()
865
+ injb.append(ib)
866
+
867
+ # ═══ LLM (16L, causal, native GQA) ═══
868
+ llm_h = buf(Se, 2048)
869
+ llm_h.copy_(aux["llm_input_embeds"].to(dev).half().reshape(Se, 2048))
870
+ lw = {k: [] for k in (
871
+ "in_ln_w", "post_ln_w", "q_norm_w", "k_norm_w", "q_w", "k_w",
872
+ "v_w", "o_w", "gate_w", "up_w", "down_w",
873
+ "q_ws", "k_ws", "v_ws", "o_ws", "gate_ws", "up_ws", "down_ws")}
874
+ lw["cos"] = self._mrope_cos.data_ptr()
875
+ lw["sin"] = self._mrope_sin.data_ptr()
876
+ lw["deepstack_inject"] = inject
877
+ for li in range(16):
878
+ qkv = self._llm_qkv_w[li] # fp8 (2048, 4096) [K, NHQ·HD+2·NHKV·HD]
879
+ q = K(qkv[:, :2048].contiguous())
880
+ kk = K(qkv[:, 2048:3072].contiguous())
881
+ v = K(qkv[:, 3072:4096].contiguous())
882
+ qkv_ws = wsc(self._llm_alpha[li * 5 + 0])
883
+ lw["in_ln_w"].append(self._llm_input_ln_w[li].data_ptr())
884
+ lw["post_ln_w"].append(self._llm_post_ln_w[li].data_ptr())
885
+ lw["q_norm_w"].append(self._llm_q_norm_w[li].data_ptr())
886
+ lw["k_norm_w"].append(self._llm_k_norm_w[li].data_ptr())
887
+ lw["q_w"].append(q.data_ptr())
888
+ lw["k_w"].append(kk.data_ptr())
889
+ lw["v_w"].append(v.data_ptr())
890
+ lw["q_ws"].append(qkv_ws)
891
+ lw["k_ws"].append(qkv_ws)
892
+ lw["v_ws"].append(qkv_ws)
893
+ lw["o_w"].append(self._llm_o_w[li].data_ptr())
894
+ lw["o_ws"].append(wsc(self._llm_alpha[li * 5 + 1]))
895
+ lw["gate_w"].append(self._llm_gate_w[li].data_ptr())
896
+ lw["gate_ws"].append(wsc(self._llm_alpha[li * 5 + 2]))
897
+ lw["up_w"].append(self._llm_up_w[li].data_ptr())
898
+ lw["up_ws"].append(wsc(self._llm_alpha[li * 5 + 3]))
899
+ lw["down_w"].append(self._llm_down_w[li].data_ptr())
900
+ lw["down_ws"].append(wsc(self._llm_alpha[li * 5 + 4]))
901
+ llm_scales = {
902
+ "act_qkv": adv(self._llm_act_qkv_dev),
903
+ "act_o": adv(self._llm_act_o_dev),
904
+ "act_gateup": adv(self._llm_act_gateup_dev),
905
+ "act_down": adv(self._llm_act_down_dev)}
906
+ # AMD delta vs the RTX mixin: the Q/K/V descale GEMMs land in the
907
+ # backend slots (Q 16 heads, K/V native 8 heads) and aiter runs
908
+ # GQA internally, so the RTX Q/K/V staging buffers and the
909
+ # K_exp/V_exp expand scratch are not allocated at all.
910
+ llm_bufs = {
911
+ "h": llm_h.data_ptr(), "xn": buf(Se, 2048).data_ptr(),
912
+ "xn_fp8": buf8(Se, 2048).data_ptr(),
913
+ "o_proj_out": buf(Se, 2048).data_ptr(),
914
+ "gate_out": buf(Se, 6144).data_ptr(),
915
+ "up_out": buf(Se, 6144).data_ptr(),
916
+ "gu_fp8": buf8(Se, 6144).data_ptr()}
917
+ llm_dims = {"S": Se, "D": 2048, "NHQ": 16, "NHKV": 8,
918
+ "HD": 128, "FF": 6144}
919
+ P.qwen3vl_llm_forward(
920
+ gemm=gemm, fvk=fvkm, bufs=llm_bufs, weights=lw,
921
+ scales_dev=llm_scales, dims=llm_dims, attn=attn,
922
+ fused_epilogue=fused)
923
+
924
+ # ═══ vlln + VL self-attn (4L) ═══
925
+ vlsa_h = buf(Se, 2048)
926
+ vlln_bufs = {"x": llm_h.data_ptr(), "out": vlsa_h.data_ptr()}
927
+ vlln_weights = {"vlln_w": self._vlln_w.data_ptr(),
928
+ "vlln_b": self._vlln_b.data_ptr()}
929
+ P.vlln_forward(
930
+ gemm=gemm, fvk=fvkm, bufs=vlln_bufs, weights=vlln_weights,
931
+ dims={"S": Se, "D": 2048})
932
+ vsw = {k: [] for k in (
933
+ "norm1_w", "norm1_b", "norm3_w", "norm3_b", "q_w", "q_b",
934
+ "k_w", "k_b", "v_w", "v_b", "o_w", "o_b", "fc1_w", "fc1_b",
935
+ "fc2_w", "fc2_b", "q_ws", "k_ws", "v_ws", "o_ws",
936
+ "fc1_ws", "fc2_ws")}
937
+ for li in range(4):
938
+ vsw["norm1_w"].append(self._vlsa_norm1_w[li].data_ptr())
939
+ vsw["norm1_b"].append(self._vlsa_norm1_b[li].data_ptr())
940
+ vsw["norm3_w"].append(self._vlsa_norm3_w[li].data_ptr())
941
+ vsw["norm3_b"].append(self._vlsa_norm3_b[li].data_ptr())
942
+ vsw["q_w"].append(self._vlsa_q_w[li].data_ptr())
943
+ vsw["q_b"].append(self._vlsa_q_b[li].data_ptr())
944
+ vsw["q_ws"].append(wsc(self._vlsa_alpha[li * 6 + 0]))
945
+ vsw["k_w"].append(self._vlsa_k_w[li].data_ptr())
946
+ vsw["k_b"].append(self._vlsa_k_b[li].data_ptr())
947
+ vsw["k_ws"].append(wsc(self._vlsa_alpha[li * 6 + 1]))
948
+ vsw["v_w"].append(self._vlsa_v_w[li].data_ptr())
949
+ vsw["v_b"].append(self._vlsa_v_b[li].data_ptr())
950
+ vsw["v_ws"].append(wsc(self._vlsa_alpha[li * 6 + 2]))
951
+ vsw["o_w"].append(self._vlsa_o_w[li].data_ptr())
952
+ vsw["o_b"].append(self._vlsa_o_b[li].data_ptr())
953
+ vsw["o_ws"].append(wsc(self._vlsa_alpha[li * 6 + 3]))
954
+ vsw["fc1_w"].append(self._vlsa_fc1_w[li].data_ptr())
955
+ vsw["fc1_b"].append(self._vlsa_fc1_b[li].data_ptr())
956
+ vsw["fc1_ws"].append(wsc(self._vlsa_alpha[li * 6 + 4]))
957
+ vsw["fc2_w"].append(self._vlsa_fc2_w[li].data_ptr())
958
+ vsw["fc2_b"].append(self._vlsa_fc2_b[li].data_ptr())
959
+ vsw["fc2_ws"].append(wsc(self._vlsa_alpha[li * 6 + 5]))
960
+ vlsa_scales = {
961
+ "act_qkv": adv(self._vlsa_act_qkv_dev),
962
+ "act_o": adv(self._vlsa_act_o_dev),
963
+ "act_fc1": adv(self._vlsa_act_fc1_dev),
964
+ "act_fc2": adv(self._vlsa_act_fc2_dev)}
965
+ vlsa_bufs = {"h": vlsa_h.data_ptr(), "xn": buf(Se, 2048).data_ptr(),
966
+ "xn_fp8": buf8(Se, 2048).data_ptr(),
967
+ "o_proj_out": buf(Se, 2048).data_ptr(),
968
+ "fc1_out": buf(Se, 8192).data_ptr(),
969
+ "fc1_fp8": buf8(Se, 8192).data_ptr()}
970
+ vlsa_dims = {"T": Se, "D": 2048, "NH": 32, "HD": 64,
971
+ "ff_inner": 8192}
972
+ P.vl_self_attn_forward(
973
+ gemm=gemm, fvk=fvkm, bufs=vlsa_bufs,
974
+ weights=vsw, scales_dev=vlsa_scales, dims=vlsa_dims, attn=attn,
975
+ fused_epilogue=fused, alphas=vlsa_alphas)
976
+ torch.cuda.synchronize()
977
+
978
+ vit_dims = {"S": Sv, "D": 1024, "NH": 16, "HD": 64,
979
+ "ff_inner": 4096, "Sper_view": Sv // nv}
980
+ vlln_dims = {"S": Se, "D": 2048}
981
+
982
+ def _kbb_forward(stream=0):
983
+ scell[0] = stream
984
+ P.qwen3vl_vit_forward(
985
+ gemm=gemm, fvk=fvkm, bufs=vit_bufs, weights=vw,
986
+ scales_dev=vit_scales, dims=vit_dims, attn=attn,
987
+ deepstack_taps=tap_layers, deepstack_capture=dcap,
988
+ stream=stream, fused_epilogue=fused, alphas=vit_alphas)
989
+ P.deepstack_merge_forward(
990
+ gemm=gemm, fvk=fvkm, bufs=ds_bufs, weights=dsw,
991
+ scales_dev=ds_scales, dims=ds_dims, stream=stream,
992
+ fused_epilogue=fused, alphas=ds_alphas)
993
+ for j in range(3):
994
+ injb[j].zero_()
995
+ injb[j].index_copy_(0, vis_idx, ds_out[j])
996
+ P.qwen3vl_llm_forward(
997
+ gemm=gemm, fvk=fvkm, bufs=llm_bufs, weights=lw,
998
+ scales_dev=llm_scales, dims=llm_dims, attn=attn,
999
+ stream=stream, fused_epilogue=fused)
1000
+ P.vlln_forward(
1001
+ gemm=gemm, fvk=fvkm, bufs=vlln_bufs,
1002
+ weights=vlln_weights, dims=vlln_dims, stream=stream)
1003
+ P.vl_self_attn_forward(
1004
+ gemm=gemm, fvk=fvkm, bufs=vlsa_bufs, weights=vsw,
1005
+ scales_dev=vlsa_scales, dims=vlsa_dims, attn=attn,
1006
+ stream=stream, fused_epilogue=fused, alphas=vlsa_alphas)
1007
+ return vlsa_h
1008
+
1009
+ self._kbb_forward = _kbb_forward
1010
+ self._kbb_vit_h = vit_h
1011
+ self._kbb_llm_h = llm_h
1012
+ self._kbb_vlsa_h = vlsa_h
1013
+ return vlsa_h.unsqueeze(0)