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,17 @@
1
+ """FlashRT -- Nex-N2-mini model pipelines.
2
+
3
+ pipeline_rtx.py - Nexn2Dims (static dims) + Nexn2Pipeline (BF16 HF
4
+ reference, used only for kernelized=False).
5
+
6
+ Nex-N2-mini is the MoE sibling of the dense qwen36 family
7
+ (architectures=Qwen3_5MoeForConditionalGeneration / model_type=qwen3_5_moe):
8
+ hybrid Gated-DeltaNet + softmax-attention with a fine-grained 256-expert MoE
9
+ FFN. The production NVFP4 kernel forward/decode (CUDA-graph, chunked
10
+ long-context prefill) lives in flash_rt.frontends.torch._nexn2_rtx_*.
11
+ """
12
+
13
+ from flash_rt.models.nexn2.pipeline_rtx import Nexn2Pipeline
14
+
15
+ __all__ = [
16
+ 'Nexn2Pipeline',
17
+ ]
@@ -0,0 +1,137 @@
1
+ """FlashRT -- Nex-N2-mini static dims + BF16 reference pipeline.
2
+
3
+ This module holds the static dimension constants (``Nexn2Dims``) and the
4
+ ``Nexn2Pipeline`` BF16-eager wrapper around the HF reference model. The
5
+ reference pipeline is used **only** when the frontend is constructed with
6
+ ``kernelized=False`` (the correctness baseline for the golden cosine fixture);
7
+ the production path -- NVFP4 kernels, CUDA-graph decode, chunked long-context
8
+ prefill -- lives in the frontend forward/decode modules
9
+ (``flash_rt.frontends.torch._nexn2_rtx_{forward,decode}``), not here.
10
+
11
+ Architecture summary (Nex-N2-mini = model_type qwen3_5_moe)::
12
+
13
+ [input_ids]
14
+ |
15
+ v embed_tokens (BF16, vocab=248320, hidden=2048)
16
+ v
17
+ 40 decoder layers, alternating linear-attn (3) + full-attn (1):
18
+ layer 0,1,2: linear_attention (Gated DeltaNet, conv1d k=4,
19
+ 16 K-heads / 32 V-heads)
20
+ layer 3: full_attention (GQA 16Q/2KV, head_dim=256,
21
+ output_gate, partial RoPE 0.25)
22
+ layer 4..39: same pattern repeats (linear x3, full x1) ...
23
+ |
24
+ v per layer: RMSNorm -> attn (linear or full)
25
+ v + residual -> RMSNorm -> MoE FFN -> residual
26
+ v MoE: 256 experts, top-8 routed + 1 shared expert
27
+ v
28
+ v final RMSNorm -> lm_head (BF16, untied)
29
+ v
30
+ [logits: (B, S, 248320)]
31
+
32
+ The config declares one MTP (multi-token-prediction) layer, but the released
33
+ Nex-N2-mini checkpoint ships no MTP tensors, so speculative decode is not wired.
34
+ """
35
+
36
+ from __future__ import annotations
37
+
38
+ from dataclasses import dataclass, field
39
+ from typing import Any
40
+
41
+
42
+ @dataclass
43
+ class Nexn2Dims:
44
+ """Static dimension constants for Nex-N2-mini.
45
+
46
+ Source: config.json:text_config (model_type=qwen3_5_moe_text). Fixed
47
+ for the mini (35B-A3B) variant; if another size is added later this
48
+ becomes a per-checkpoint loader instead of a class-level constant.
49
+ """
50
+ hidden: int = 2048
51
+ num_layers: int = 40
52
+ full_attn_period: int = 4 # full at indices 3, 7, ..., 39
53
+ vocab_size: int = 248320
54
+ rms_norm_eps: float = 1e-6
55
+
56
+ # full-attention sites (10 layers)
57
+ full_q_heads: int = 16
58
+ full_kv_heads: int = 2 # GQA 8:1
59
+ full_head_dim: int = 256
60
+ partial_rotary_factor: float = 0.25 # rotary_dim = 64
61
+ rope_theta: float = 1.0e7
62
+ mrope_section: tuple[int, ...] = (11, 11, 10)
63
+
64
+ # linear-attention sites (30 layers, Gated DeltaNet)
65
+ lin_k_heads: int = 16
66
+ lin_v_heads: int = 32 # differs from qwen36 (48)
67
+ lin_head_dim: int = 128
68
+ lin_conv_kernel: int = 4
69
+
70
+ # MoE FFN (every layer)
71
+ moe_num_experts: int = 256
72
+ moe_experts_per_tok: int = 8
73
+ moe_intermediate: int = 512
74
+ shared_expert_intermediate: int = 512
75
+
76
+ # MTP head
77
+ mtp_layers: int = 1
78
+
79
+
80
+ class Nexn2Pipeline:
81
+ """BF16 HF reference pipeline for Nex-N2-mini.
82
+
83
+ Hosts an HF reference model and delegates ``forward`` / ``generate``. Used
84
+ only by the frontend's ``kernelized=False`` path (the correctness baseline);
85
+ the production NVFP4 kernel forward/decode lives in
86
+ ``flash_rt.frontends.torch._nexn2_rtx_{forward,decode}``.
87
+ """
88
+
89
+ DIMS = Nexn2Dims()
90
+
91
+ def __init__(self, hf_model: Any) -> None:
92
+ """Wrap an HF model object (the qwen3_5_moe auto-loader output)."""
93
+ self.hf = hf_model
94
+ self.config = hf_model.config
95
+ text_cfg = getattr(self.config, 'text_config', self.config)
96
+ # Sanity-check the dim assumptions against the checkpoint config.
97
+ assert text_cfg.hidden_size == self.DIMS.hidden, (
98
+ f'expected hidden={self.DIMS.hidden}, got {text_cfg.hidden_size}'
99
+ )
100
+ assert text_cfg.num_hidden_layers == self.DIMS.num_layers
101
+ assert text_cfg.head_dim == self.DIMS.full_head_dim
102
+ assert text_cfg.num_experts == self.DIMS.moe_num_experts
103
+ assert (
104
+ text_cfg.layer_types.count('full_attention')
105
+ == self.DIMS.num_layers // self.DIMS.full_attn_period
106
+ )
107
+
108
+ def forward(self, input_ids):
109
+ """Single forward pass: token IDs -> logits (BF16 HF reference).
110
+
111
+ Args:
112
+ input_ids: (B, S) torch.long on cuda.
113
+
114
+ Returns:
115
+ logits: (B, S, vocab_size) bf16 on cuda.
116
+ """
117
+ import torch # local import; pipeline_rtx is import-time-light.
118
+ with torch.no_grad():
119
+ out = self.hf(
120
+ input_ids=input_ids, use_cache=False, return_dict=True,
121
+ )
122
+ return out.logits
123
+
124
+ def generate(self, input_ids, *, max_new_tokens: int, do_sample: bool = False):
125
+ """Greedy/sampled autoregressive generate (BF16 HF reference path).
126
+
127
+ The production decode (CUDA-graph, on-device argmax) is in the
128
+ frontend's kernelized path, not here.
129
+ """
130
+ import torch
131
+ with torch.no_grad():
132
+ return self.hf.generate(
133
+ input_ids=input_ids,
134
+ max_new_tokens=max_new_tokens,
135
+ do_sample=do_sample,
136
+ use_cache=True,
137
+ )
@@ -0,0 +1,30 @@
1
+ """FlashRT — OmniVoice TTS model pipeline.
2
+
3
+ Mixed BF16 CFG + FP4 noCFG acceleration preserves audio quality at 5.0x
4
+ throughput on Blackwell SM120 GPUs.
5
+
6
+ Per the unified API contract:
7
+ inject() — patch OmniVoice model for FlashRT acceleration
8
+ free_encoder() — release encoder weights (~600 MB VRAM saved)
9
+ eject() — restore original forward and generate methods
10
+
11
+ See docs/PERFORMANCE_OMNIVOICE.md for performance specifications.
12
+ """
13
+
14
+ from flash_rt.models.omnivoice.pipeline_rtx import (
15
+ FlashRTLlm,
16
+ FlashRTLlmBF16,
17
+ inject,
18
+ free_encoder,
19
+ eject,
20
+ _check_kernels,
21
+ )
22
+
23
+ __all__ = [
24
+ "FlashRTLlm",
25
+ "FlashRTLlmBF16",
26
+ "inject",
27
+ "free_encoder",
28
+ "eject",
29
+ "_check_kernels",
30
+ ]