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.
- flash_rt/__init__.py +107 -0
- flash_rt/_extensions.py +119 -0
- flash_rt/amd/__init__.py +9 -0
- flash_rt/amd/core/__init__.py +0 -0
- flash_rt/amd/core/hip_buffer.py +176 -0
- flash_rt/amd/core/hip_graph.py +102 -0
- flash_rt/amd/frontends/__init__.py +1 -0
- flash_rt/amd/frontends/torch/__init__.py +1 -0
- flash_rt/amd/frontends/torch/groot_n17.py +1013 -0
- flash_rt/amd/frontends/torch/pi05.py +1431 -0
- flash_rt/amd/hardware/__init__.py +1 -0
- flash_rt/amd/hardware/cdna4/__init__.py +1 -0
- flash_rt/amd/hardware/cdna4/attn_backend.py +338 -0
- flash_rt/amd/hardware/cdna4/attn_backend_aiter.py +410 -0
- flash_rt/amd/hardware/cdna4/attn_backend_groot_n17.py +297 -0
- flash_rt/amd/models/__init__.py +1 -0
- flash_rt/amd/models/groot_n17/__init__.py +1 -0
- flash_rt/amd/models/groot_n17/pipeline.py +1058 -0
- flash_rt/amd/models/pi05/__init__.py +1 -0
- flash_rt/amd/models/pi05/pipeline.py +1860 -0
- flash_rt/api.py +1144 -0
- flash_rt/catalog/__init__.py +39 -0
- flash_rt/catalog/binding.py +412 -0
- flash_rt/catalog/bindings/cosmos3_video_pipeline.yaml +94 -0
- flash_rt/catalog/bindings/groot_n16_dit.yaml +23 -0
- flash_rt/catalog/bindings/groot_n16_llm.yaml +21 -0
- flash_rt/catalog/bindings/groot_n16_pipeline.yaml +107 -0
- flash_rt/catalog/bindings/groot_n16_tick.yaml +26 -0
- flash_rt/catalog/bindings/groot_n16_vision.yaml +23 -0
- flash_rt/catalog/bindings/groot_n17_pipeline.yaml +117 -0
- flash_rt/catalog/bindings/lingbot_vla_pipeline.yaml +106 -0
- flash_rt/catalog/bindings/motus_tick.yaml +120 -0
- flash_rt/catalog/bindings/nexn2_pipeline.yaml +112 -0
- flash_rt/catalog/bindings/pi05.yaml +29 -0
- flash_rt/catalog/bindings/pi05_prefix.yaml +21 -0
- flash_rt/catalog/bindings/pi05_tick.yaml +94 -0
- flash_rt/catalog/bindings/pi05_vision.yaml +23 -0
- flash_rt/catalog/bindings/qwen25_15b.yaml +22 -0
- flash_rt/catalog/bindings/qwen36_27b_pipeline.yaml +130 -0
- flash_rt/catalog/bindings/qwen3_8b.yaml +22 -0
- flash_rt/catalog/bindings/qwen3_8b_pipeline.yaml +98 -0
- flash_rt/catalog/bindings/qwen3_vl_8b_pipeline.yaml +131 -0
- flash_rt/catalog/bindings/qwen3_vl_8b_text.yaml +22 -0
- flash_rt/catalog/bindings/qwen3_vl_8b_vision.yaml +23 -0
- flash_rt/catalog/bindings/smolvla_base.yaml +21 -0
- flash_rt/catalog/bindings/smolvla_expert.yaml +21 -0
- flash_rt/catalog/bindings/smolvla_pipeline.yaml +93 -0
- flash_rt/catalog/bindings/smolvla_tick.yaml +27 -0
- flash_rt/catalog/bindings/smolvla_vision.yaml +23 -0
- flash_rt/catalog/bindings/wan22_video_pipeline.yaml +119 -0
- flash_rt/catalog/registry.py +108 -0
- flash_rt/catalog/structures/__init__.py +0 -0
- flash_rt/catalog/structures/adaln_producer/__init__.py +0 -0
- flash_rt/catalog/structures/adaln_producer/reference.py +40 -0
- flash_rt/catalog/structures/adaln_producer/structure.yaml +106 -0
- flash_rt/catalog/structures/attention_core/__init__.py +0 -0
- flash_rt/catalog/structures/attention_core/reference.py +36 -0
- flash_rt/catalog/structures/attention_core/structure.yaml +72 -0
- flash_rt/catalog/structures/autoregressive_decode_pipeline/structure.yaml +82 -0
- flash_rt/catalog/structures/cadence_static/__init__.py +0 -0
- flash_rt/catalog/structures/cadence_static/reference.py +24 -0
- flash_rt/catalog/structures/cadence_static/structure.yaml +56 -0
- flash_rt/catalog/structures/decoder_block/__init__.py +0 -0
- flash_rt/catalog/structures/decoder_block/reference.py +38 -0
- flash_rt/catalog/structures/decoder_block/structure.yaml +86 -0
- flash_rt/catalog/structures/decoder_ffn/__init__.py +0 -0
- flash_rt/catalog/structures/decoder_ffn/reference.py +64 -0
- flash_rt/catalog/structures/decoder_ffn/structure.yaml +44 -0
- flash_rt/catalog/structures/gated_delta_core/reference.py +54 -0
- flash_rt/catalog/structures/gated_delta_core/structure.yaml +60 -0
- flash_rt/catalog/structures/linear_proj/__init__.py +3 -0
- flash_rt/catalog/structures/linear_proj/reference.py +35 -0
- flash_rt/catalog/structures/linear_proj/structure.yaml +74 -0
- flash_rt/catalog/structures/modnorm_qkv_chain/__init__.py +1 -0
- flash_rt/catalog/structures/modnorm_qkv_chain/reference.py +39 -0
- flash_rt/catalog/structures/modnorm_qkv_chain/structure.yaml +59 -0
- flash_rt/catalog/structures/norm_fused/__init__.py +0 -0
- flash_rt/catalog/structures/norm_fused/reference.py +26 -0
- flash_rt/catalog/structures/norm_fused/structure.yaml +50 -0
- flash_rt/catalog/structures/patch_projection/reference.py +15 -0
- flash_rt/catalog/structures/patch_projection/structure.yaml +53 -0
- flash_rt/catalog/structures/qk_norm_rope/__init__.py +3 -0
- flash_rt/catalog/structures/qk_norm_rope/reference.py +114 -0
- flash_rt/catalog/structures/qk_norm_rope/structure.yaml +83 -0
- flash_rt/catalog/structures/qkv_pack/__init__.py +0 -0
- flash_rt/catalog/structures/qkv_pack/reference.py +33 -0
- flash_rt/catalog/structures/qkv_pack/structure.yaml +65 -0
- flash_rt/catalog/structures/qkv_rope/__init__.py +1 -0
- flash_rt/catalog/structures/qkv_rope/reference.py +39 -0
- flash_rt/catalog/structures/qkv_rope/structure.yaml +55 -0
- flash_rt/catalog/structures/video_generation_pipeline/__init__.py +2 -0
- flash_rt/catalog/structures/video_generation_pipeline/structure.yaml +81 -0
- flash_rt/catalog/structures/vision_ffn/__init__.py +0 -0
- flash_rt/catalog/structures/vision_ffn/reference.py +35 -0
- flash_rt/catalog/structures/vision_ffn/structure.yaml +42 -0
- flash_rt/catalog/structures/vla_tick_pipeline/__init__.py +7 -0
- flash_rt/catalog/structures/vla_tick_pipeline/structure.yaml +82 -0
- flash_rt/configs/__init__.py +0 -0
- flash_rt/configs/cosmos3_edge.yaml +21 -0
- flash_rt/configs/cosmos3_video.yaml +24 -0
- flash_rt/configs/groot.yaml +73 -0
- flash_rt/configs/groot_n17.yaml +53 -0
- flash_rt/configs/hyvla.yaml +65 -0
- flash_rt/configs/ltx25.yaml +41 -0
- flash_rt/configs/motus.yaml +85 -0
- flash_rt/configs/nexn2.yaml +79 -0
- flash_rt/configs/pi0.yaml +38 -0
- flash_rt/configs/pi05.yaml +38 -0
- flash_rt/configs/qwen36.yaml +68 -0
- flash_rt/configs/wan22_ti2v_5b.yaml +24 -0
- flash_rt/core/__init__.py +0 -0
- flash_rt/core/calibration.py +301 -0
- flash_rt/core/calibration_api.py +70 -0
- flash_rt/core/config.py +96 -0
- flash_rt/core/context.py +47 -0
- flash_rt/core/cuda_buffer.py +189 -0
- flash_rt/core/cuda_graph.py +81 -0
- flash_rt/core/parity.py +37 -0
- flash_rt/core/precision_spec.py +164 -0
- flash_rt/core/quant/__init__.py +0 -0
- flash_rt/core/quant/calibrator.py +170 -0
- flash_rt/core/quantization.py +73 -0
- flash_rt/core/rl/__init__.py +75 -0
- flash_rt/core/rl/acp_tags.py +51 -0
- flash_rt/core/rl/advantage.py +163 -0
- flash_rt/core/rl/cfg_sampler.py +72 -0
- flash_rt/core/rl/reward.py +233 -0
- flash_rt/core/rl/value_function.py +198 -0
- flash_rt/core/thor_frontend_utils.py +152 -0
- flash_rt/core/utils/__init__.py +0 -0
- flash_rt/core/utils/actions.py +19 -0
- flash_rt/core/utils/hardware.py +50 -0
- flash_rt/core/utils/norm_stats.py +359 -0
- flash_rt/core/utils/pi05_prompt.py +35 -0
- flash_rt/core/weights/__init__.py +0 -0
- flash_rt/core/weights/loader.py +135 -0
- flash_rt/core/weights/transformer.py +691 -0
- flash_rt/core/weights/weight_cache.py +147 -0
- flash_rt/datasets/__init__.py +11 -0
- flash_rt/datasets/libero.py +306 -0
- flash_rt/executors/__init__.py +6 -0
- flash_rt/executors/fp4_utils.py +241 -0
- flash_rt/executors/fp4_utils_cb.py +207 -0
- flash_rt/executors/jax_weights.py +270 -0
- flash_rt/executors/torch_weights.py +500 -0
- flash_rt/executors/weight_loader.py +331 -0
- flash_rt/frontends/__init__.py +8 -0
- flash_rt/frontends/_fp8_layout.py +32 -0
- flash_rt/frontends/jax/__init__.py +1 -0
- flash_rt/frontends/jax/_pi05_thor_spec.py +52 -0
- flash_rt/frontends/jax/_pi0_thor_spec.py +32 -0
- flash_rt/frontends/jax/_thor_spec_common.py +114 -0
- flash_rt/frontends/jax/pi05_rtx.py +576 -0
- flash_rt/frontends/jax/pi05_thor.py +2768 -0
- flash_rt/frontends/jax/pi05_thor_fp4.py +879 -0
- flash_rt/frontends/jax/pi0_rtx.py +483 -0
- flash_rt/frontends/jax/pi0_thor.py +1425 -0
- flash_rt/frontends/jax/pi0fast.py +1337 -0
- flash_rt/frontends/jetson_pi/__init__.py +12 -0
- flash_rt/frontends/jetson_pi/llm.py +261 -0
- flash_rt/frontends/jetson_pi/mllm.py +289 -0
- flash_rt/frontends/jetson_pi/pi0.py +420 -0
- flash_rt/frontends/torch/__init__.py +1 -0
- flash_rt/frontends/torch/_chameleon_quant.py +251 -0
- flash_rt/frontends/torch/_chameleon_rtx_sm87_spec.py +85 -0
- flash_rt/frontends/torch/_chameleon_thor_spec.py +92 -0
- flash_rt/frontends/torch/_cosmos3_edge_thor_spec.py +177 -0
- flash_rt/frontends/torch/_groot_n17_rtx_spec.py +13 -0
- flash_rt/frontends/torch/_groot_n17_thor_spec.py +406 -0
- flash_rt/frontends/torch/_groot_thor_spec.py +105 -0
- flash_rt/frontends/torch/_higgs_audio_v3_bf16.py +374 -0
- flash_rt/frontends/torch/_higgs_audio_v3_fp8.py +464 -0
- flash_rt/frontends/torch/_hyvla_thor_spec.py +196 -0
- flash_rt/frontends/torch/_lingbot_thor_spec.py +321 -0
- flash_rt/frontends/torch/_motus_rtx_spec.py +47 -0
- flash_rt/frontends/torch/_nexn2_rtx_decode.py +1621 -0
- flash_rt/frontends/torch/_nexn2_rtx_forward.py +1815 -0
- flash_rt/frontends/torch/_nexn2_rtx_nvfp4_weights.py +416 -0
- flash_rt/frontends/torch/_pi05_thor_spec.py +100 -0
- flash_rt/frontends/torch/_pi0_thor_spec.py +81 -0
- flash_rt/frontends/torch/_qwen36_rtx_dflash_forward.py +940 -0
- flash_rt/frontends/torch/_qwen36_rtx_dflash_weights.py +396 -0
- flash_rt/frontends/torch/_qwen36_rtx_nvfp4_weights.py +741 -0
- flash_rt/frontends/torch/_qwen36_rtx_turboquant.py +872 -0
- flash_rt/frontends/torch/_qwen36_rtx_weights.py +411 -0
- flash_rt/frontends/torch/_qwen3_rtx_nvfp4_weights.py +575 -0
- flash_rt/frontends/torch/_qwen3_vl_bf16_weights.py +191 -0
- flash_rt/frontends/torch/_qwen3_vl_fp8_weights.py +245 -0
- flash_rt/frontends/torch/_qwen3_vl_geometry.py +337 -0
- flash_rt/frontends/torch/_qwen3_vl_vision_rtx.py +625 -0
- flash_rt/frontends/torch/_template/attention.py +124 -0
- flash_rt/frontends/torch/_template/frontend.py +330 -0
- flash_rt/frontends/torch/_template/pipeline.py +263 -0
- flash_rt/frontends/torch/_template/weights_spec.py +215 -0
- flash_rt/frontends/torch/_thor_spec_common.py +147 -0
- flash_rt/frontends/torch/chameleon_rtx_sm87.py +721 -0
- flash_rt/frontends/torch/chameleon_thor.py +911 -0
- flash_rt/frontends/torch/cosmos3_edge_thor.py +571 -0
- flash_rt/frontends/torch/cosmos3_video_rtx.py +130 -0
- flash_rt/frontends/torch/groot_n17_rtx.py +152 -0
- flash_rt/frontends/torch/groot_n17_rtx_fp16.py +655 -0
- flash_rt/frontends/torch/groot_n17_rtx_fp8.py +582 -0
- flash_rt/frontends/torch/groot_n17_rtx_sm89.py +609 -0
- flash_rt/frontends/torch/groot_n17_rtx_sm89_fp16.py +652 -0
- flash_rt/frontends/torch/groot_n17_thor.py +1965 -0
- flash_rt/frontends/torch/groot_n17_thor_fp16.py +49 -0
- flash_rt/frontends/torch/groot_n17_thor_fp4.py +165 -0
- flash_rt/frontends/torch/groot_n17_thor_fp8.py +780 -0
- flash_rt/frontends/torch/groot_rtx.py +1876 -0
- flash_rt/frontends/torch/groot_rtx_fp16.py +1162 -0
- flash_rt/frontends/torch/groot_thor.py +3623 -0
- flash_rt/frontends/torch/groot_thor_fp16.py +28 -0
- flash_rt/frontends/torch/higgs_audio_v3_rtx.py +601 -0
- flash_rt/frontends/torch/hyvla_orin.py +306 -0
- flash_rt/frontends/torch/hyvla_thor.py +683 -0
- flash_rt/frontends/torch/lingbot_thor.py +116 -0
- flash_rt/frontends/torch/ltx25_rtx.py +378 -0
- flash_rt/frontends/torch/motus_rtx.py +1562 -0
- flash_rt/frontends/torch/nexn2_rtx.py +310 -0
- flash_rt/frontends/torch/pi05_rtx.py +1949 -0
- flash_rt/frontends/torch/pi05_rtx_fp16.py +1806 -0
- flash_rt/frontends/torch/pi05_thor.py +3083 -0
- flash_rt/frontends/torch/pi05_thor_fp4.py +1500 -0
- flash_rt/frontends/torch/pi0_rtx.py +957 -0
- flash_rt/frontends/torch/pi0_thor.py +1405 -0
- flash_rt/frontends/torch/pi0fast.py +1402 -0
- flash_rt/frontends/torch/qwen36_moe.py +426 -0
- flash_rt/frontends/torch/qwen36_moe_rtx.py +20 -0
- flash_rt/frontends/torch/qwen36_rtx.py +12388 -0
- flash_rt/frontends/torch/qwen36_spark.py +200 -0
- flash_rt/frontends/torch/qwen36_thor.py +1320 -0
- flash_rt/frontends/torch/qwen3_rtx.py +2166 -0
- flash_rt/frontends/torch/qwen3_vl_fp8_sm89.py +912 -0
- flash_rt/frontends/torch/qwen3_vl_fp8_sm89_multimodal.py +456 -0
- flash_rt/frontends/torch/qwen3_vl_rtx.py +642 -0
- flash_rt/frontends/torch/qwen3_vl_rtx_bf16.py +1031 -0
- flash_rt/frontends/torch/qwen3_vl_thor.py +856 -0
- flash_rt/frontends/torch/wan22_rtx.py +477 -0
- flash_rt/hardware/__init__.py +311 -0
- flash_rt/hardware/backend.py +407 -0
- flash_rt/hardware/blackwell/__init__.py +12 -0
- flash_rt/hardware/rtx/__init__.py +37 -0
- flash_rt/hardware/rtx/attn_backend.py +856 -0
- flash_rt/hardware/rtx/attn_backend_batched_pi05.py +303 -0
- flash_rt/hardware/rtx/attn_backend_chameleon.py +237 -0
- flash_rt/hardware/rtx/attn_backend_groot.py +447 -0
- flash_rt/hardware/rtx/attn_backend_groot_n17.py +251 -0
- flash_rt/hardware/rtx/attn_backend_groot_n17_backbone.py +191 -0
- flash_rt/hardware/rtx/attn_backend_motus.py +128 -0
- flash_rt/hardware/rtx/attn_backend_nexn2.py +409 -0
- flash_rt/hardware/rtx/attn_backend_qwen3.py +520 -0
- flash_rt/hardware/rtx/attn_backend_qwen36.py +272 -0
- flash_rt/hardware/thor/__init__.py +9 -0
- flash_rt/hardware/thor/attn_backend.py +559 -0
- flash_rt/hardware/thor/attn_backend_chameleon.py +362 -0
- flash_rt/hardware/thor/attn_backend_groot.py +328 -0
- flash_rt/hardware/thor/attn_backend_groot_n17.py +423 -0
- flash_rt/hardware/thor/attn_backend_qwen3.py +229 -0
- flash_rt/hardware/thor/attn_backend_qwen36.py +530 -0
- flash_rt/hardware/thor/fa4_backend.py +117 -0
- flash_rt/hardware/thor/shared_primitives.py +727 -0
- flash_rt/hardware/thor/shared_primitives_batched.py +178 -0
- flash_rt/hardware/thor/shared_primitives_fp4.py +512 -0
- flash_rt/hardware/thor/vqgan_trt_backend.py +187 -0
- flash_rt/models/__init__.py +11 -0
- flash_rt/models/chameleon/__init__.py +18 -0
- flash_rt/models/chameleon/pipeline_rtx.py +305 -0
- flash_rt/models/chameleon/pipeline_thor.py +1126 -0
- flash_rt/models/chameleon/vqvae_hf.py +124 -0
- flash_rt/models/cosmos3_edge/__init__.py +37 -0
- flash_rt/models/cosmos3_edge/action_only_official.py +3447 -0
- flash_rt/models/cosmos3_edge/boundary_dump.py +151 -0
- flash_rt/models/cosmos3_edge/denoise_ref.py +325 -0
- flash_rt/models/cosmos3_edge/dump_replay.py +195 -0
- flash_rt/models/cosmos3_edge/layer_ref.py +1040 -0
- flash_rt/models/cosmos3_edge/pipeline_thor.py +549 -0
- flash_rt/models/cosmos3_edge/static_engine.py +346 -0
- flash_rt/models/cosmos3_edge/static_unipc.py +234 -0
- flash_rt/models/cosmos3_edge/vae_native.py +304 -0
- flash_rt/models/cosmos3_edge/weights.py +91 -0
- flash_rt/models/cosmos3_reasoner/__init__.py +1 -0
- flash_rt/models/cosmos3_reasoner/pipeline_thor.py +691 -0
- flash_rt/models/cosmos3_video/__init__.py +6 -0
- flash_rt/models/cosmos3_video/fm_solvers_unipc.py +808 -0
- flash_rt/models/cosmos3_video/kernels/__init__.py +22 -0
- flash_rt/models/cosmos3_video/kernels/csrc/bindings.cpp +17 -0
- flash_rt/models/cosmos3_video/kernels/csrc/fused_qk_norm_rope.cu +61 -0
- flash_rt/models/cosmos3_video/kernels/setup.py +33 -0
- flash_rt/models/cosmos3_video/pipeline_rtx.py +234 -0
- flash_rt/models/groot/__init__.py +32 -0
- flash_rt/models/groot/embodiments.py +69 -0
- flash_rt/models/groot/pipeline_rtx.py +1179 -0
- flash_rt/models/groot/pipeline_rtx_fp16.py +1034 -0
- flash_rt/models/groot/pipeline_thor.py +981 -0
- flash_rt/models/groot_n17/__init__.py +50 -0
- flash_rt/models/groot_n17/calibration.py +439 -0
- flash_rt/models/groot_n17/embodiments.py +38 -0
- flash_rt/models/groot_n17/mrope_table.py +178 -0
- flash_rt/models/groot_n17/pipeline_rtx.py +11 -0
- flash_rt/models/groot_n17/pipeline_rtx_fp16.py +836 -0
- flash_rt/models/groot_n17/pipeline_rtx_fp8.py +413 -0
- flash_rt/models/groot_n17/pipeline_rtx_sm89.py +541 -0
- flash_rt/models/groot_n17/pipeline_thor.py +1386 -0
- flash_rt/models/higgs_audio_v3/__init__.py +14 -0
- flash_rt/models/higgs_audio_v3/_codec/__init__.py +0 -0
- flash_rt/models/higgs_audio_v3/_codec/env_guard.py +42 -0
- flash_rt/models/higgs_audio_v3/_codec/tokenizer_config.json +129 -0
- flash_rt/models/higgs_audio_v3/_codec/tokenizer_model.py +940 -0
- flash_rt/models/higgs_audio_v3/codec.py +81 -0
- flash_rt/models/higgs_audio_v3/pipeline_rtx.py +64 -0
- flash_rt/models/hyvla/__init__.py +1 -0
- flash_rt/models/hyvla/pipeline_orin.py +430 -0
- flash_rt/models/hyvla/pipeline_thor.py +572 -0
- flash_rt/models/lingbot/__init__.py +17 -0
- flash_rt/models/lingbot/_csrc_loader.py +70 -0
- flash_rt/models/lingbot/buffer_binder.py +156 -0
- flash_rt/models/lingbot/calibration.py +163 -0
- flash_rt/models/lingbot/forward.py +784 -0
- flash_rt/models/lingbot/fp4_ops.py +90 -0
- flash_rt/models/lingbot/graph_runner.py +265 -0
- flash_rt/models/lingbot/kernel_ops.py +1487 -0
- flash_rt/models/lingbot/mixed_attention.py +793 -0
- flash_rt/models/lingbot/norms.py +113 -0
- flash_rt/models/lingbot/pipeline_thor.py +169 -0
- flash_rt/models/lingbot/rope_adapter.py +156 -0
- flash_rt/models/lingbot/sample_actions.py +394 -0
- flash_rt/models/lingbot/vit.py +486 -0
- flash_rt/models/lingbot/vit_rope_adapter.py +247 -0
- flash_rt/models/ltx25/__init__.py +16 -0
- flash_rt/models/ltx25/_attn_swap.py +244 -0
- flash_rt/models/ltx25/_nvfp4_ffn_swap.py +301 -0
- flash_rt/models/ltx25/_resident_graph.py +206 -0
- flash_rt/models/melband_roformer/__init__.py +13 -0
- flash_rt/models/melband_roformer/pipeline.py +329 -0
- flash_rt/models/minimax_remover/__init__.py +27 -0
- flash_rt/models/minimax_remover/_attention.py +428 -0
- flash_rt/models/minimax_remover/_fp8_linear.py +426 -0
- flash_rt/models/minimax_remover/_fp8_manual_denoise.py +298 -0
- flash_rt/models/minimax_remover/_fp8_pipeline.py +617 -0
- flash_rt/models/minimax_remover/_kern_block.py +282 -0
- flash_rt/models/minimax_remover/_kernels.py +282 -0
- flash_rt/models/minimax_remover/_manual_denoise.py +413 -0
- flash_rt/models/minimax_remover/_nvfp4_linear.py +236 -0
- flash_rt/models/minimax_remover/_triton_flash_attn.py +139 -0
- flash_rt/models/minimax_remover/_utils.py +94 -0
- flash_rt/models/minimax_remover/_vae_nvfp4.py +569 -0
- flash_rt/models/minimax_remover/_vae_opt.py +880 -0
- flash_rt/models/minimax_remover/pipeline.py +209 -0
- flash_rt/models/motus/__init__.py +0 -0
- flash_rt/models/motus/_action_ffn_v6t_install.py +162 -0
- flash_rt/models/motus/_action_und_qkv_fp8_swap.py +266 -0
- flash_rt/models/motus/_attn_swap.py +224 -0
- flash_rt/models/motus/_awq_fp8_swap.py +525 -0
- flash_rt/models/motus/_easycache_swap.py +279 -0
- flash_rt/models/motus/_ffn_swap.py +225 -0
- flash_rt/models/motus/_fp8_swap.py +635 -0
- flash_rt/models/motus/_graph_capture.py +203 -0
- flash_rt/models/motus/_handtuned_fp8_dispatch.py +235 -0
- flash_rt/models/motus/_kv_cache_swap.py +539 -0
- flash_rt/models/motus/_linear_swap.py +248 -0
- flash_rt/models/motus/_mixcache_swap.py +276 -0
- flash_rt/models/motus/_modulate_fuse_swap.py +1649 -0
- flash_rt/models/motus/_motus_nvfp4_ffn_video_swap.py +677 -0
- flash_rt/models/motus/_norm_swap.py +240 -0
- flash_rt/models/motus/_rope_swap.py +226 -0
- flash_rt/models/motus/_stream.py +22 -0
- flash_rt/models/motus/_taylorseer_swap.py +275 -0
- flash_rt/models/motus/_teacache_swap.py +210 -0
- flash_rt/models/motus/_tinyfp8_dispatch_install.py +170 -0
- flash_rt/models/motus/_und_ffn_v5t_install.py +180 -0
- flash_rt/models/motus/_vae_fp4_swap.py +849 -0
- flash_rt/models/motus/_vae_fp8_resample_swap.py +534 -0
- flash_rt/models/motus/_vae_fp8_swap.py +1082 -0
- flash_rt/models/motus/_vae_swap.py +80 -0
- flash_rt/models/motus/_vae_time_conv_fp8_swap.py +292 -0
- flash_rt/models/motus/_wan_qkv_fuse_swap.py +905 -0
- flash_rt/models/motus/pipeline_rtx.py +1164 -0
- flash_rt/models/nexn2/__init__.py +17 -0
- flash_rt/models/nexn2/pipeline_rtx.py +137 -0
- flash_rt/models/omnivoice/__init__.py +30 -0
- flash_rt/models/omnivoice/pipeline_rtx.py +546 -0
- flash_rt/models/pi0/__init__.py +9 -0
- flash_rt/models/pi0/pipeline_rtx.py +1110 -0
- flash_rt/models/pi0/pipeline_thor.py +434 -0
- flash_rt/models/pi05/__init__.py +28 -0
- flash_rt/models/pi05/pipeline_rtx.py +2209 -0
- flash_rt/models/pi05/pipeline_rtx_batched.py +1188 -0
- flash_rt/models/pi05/pipeline_rtx_cfg.py +657 -0
- flash_rt/models/pi05/pipeline_rtx_cfg_batched.py +435 -0
- flash_rt/models/pi05/pipeline_rtx_fp16.py +2276 -0
- flash_rt/models/pi05/pipeline_thor.py +929 -0
- flash_rt/models/pi05/pipeline_thor_batched.py +346 -0
- flash_rt/models/pi05/pipeline_thor_cfg.py +238 -0
- flash_rt/models/pi05/pipeline_thor_cfg_batched.py +180 -0
- flash_rt/models/pi05/runtime_export.py +449 -0
- flash_rt/models/pi0fast/__init__.py +1 -0
- flash_rt/models/pi0fast/pipeline.py +840 -0
- flash_rt/models/qwen3/__init__.py +11 -0
- flash_rt/models/qwen3/pipeline_rtx.py +90 -0
- flash_rt/models/qwen36/__init__.py +27 -0
- flash_rt/models/qwen36/pipeline_rtx.py +159 -0
- flash_rt/models/qwen3_vl/__init__.py +19 -0
- flash_rt/models/qwen3_vl/pipeline_rtx.py +145 -0
- flash_rt/models/wan22/__init__.py +1 -0
- flash_rt/models/wan22/pipeline_rtx.py +28 -0
- flash_rt/npu/__init__.py +6 -0
- flash_rt/npu/core/__init__.py +0 -0
- flash_rt/npu/core/abi.py +81 -0
- flash_rt/npu/core/acl_runtime.py +165 -0
- flash_rt/npu/core/decode_attention.py +26 -0
- flash_rt/npu/core/decoder_int8.py +399 -0
- flash_rt/npu/core/device.py +73 -0
- flash_rt/npu/core/gu_int8.py +99 -0
- flash_rt/npu/core/linear.py +283 -0
- flash_rt/npu/core/native_kernels.py +269 -0
- flash_rt/npu/core/npu_graph.py +68 -0
- flash_rt/npu/frontends/__init__.py +0 -0
- flash_rt/npu/frontends/torch/__init__.py +0 -0
- flash_rt/npu/frontends/torch/pi05.py +438 -0
- flash_rt/npu/hardware/__init__.py +9 -0
- flash_rt/npu/models/__init__.py +1 -0
- flash_rt/npu/models/pi05/__init__.py +1 -0
- flash_rt/npu/models/pi05/attention.py +128 -0
- flash_rt/npu/models/pi05/captured.py +212 -0
- flash_rt/npu/models/pi05/fast.py +666 -0
- flash_rt/npu/models/pi05/pipeline.py +336 -0
- flash_rt/npu/models/pi05/quantization.py +235 -0
- flash_rt/npu/verify.py +110 -0
- flash_rt/py.typed +0 -0
- flash_rt/refs/__init__.py +17 -0
- flash_rt/refs/pi05_cfg_reference.py +310 -0
- flash_rt/runtime/__init__.py +47 -0
- flash_rt/runtime/cuda_libraries.py +108 -0
- flash_rt/runtime/exec.py +53 -0
- flash_rt/runtime/export.py +520 -0
- flash_rt/runtime/provider.py +113 -0
- flash_rt/runtime/rtc.py +261 -0
- flash_rt/runtime/rtc_temporal_fusion.py +545 -0
- flash_rt/runtime/vlash.py +420 -0
- flash_rt/subgraphs/__init__.py +43 -0
- flash_rt/subgraphs/capture.py +179 -0
- flash_rt/subgraphs/pi05/__init__.py +3 -0
- flash_rt/subgraphs/pi05/context_action.py +67 -0
- flash_rt/subgraphs/pi05/rtc_prefix.py +79 -0
- flash_rt/subgraphs/pi05/rtc_vjp_guided.py +96 -0
- flash_rt/subgraphs/pi05/stage_plans.py +89 -0
- flash_rt/subgraphs/pi05/vlash.py +29 -0
- flash_rt/subgraphs/stage_plan.py +214 -0
- flash_rt/utils/__init__.py +1 -0
- flash_rt/utils/paligemma_tokenizer.py +135 -0
- flash_rt-0.2.0.dist-info/METADATA +1341 -0
- flash_rt-0.2.0.dist-info/RECORD +455 -0
- flash_rt-0.2.0.dist-info/WHEEL +5 -0
- flash_rt-0.2.0.dist-info/licenses/LICENSE +202 -0
- 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)
|