+
    QV-j–Š  ã                  ód  a € 0 t $ R t^ RIHt ^ RIt^ RIHt ^ RIHt ^ RI	t	^RI
Ht ^RIHtHtHtHt ^RIHt ^RIHt ]P*                  ! ]4      t]! R	R
7       ! R R4      4       t]P2                  R3R R ll4       t]	P6                  P8                  R3R R ll4       tR3R R llt]P2                  R R l4       t] ! 4       t!R]"R&   R R lt#R R lt$R4R R llt%R R lt&R R  lt'R! R" lt(R# R$ lt)R% R& lt*R' R( lt+RR]	PX                  R3R) R* llt-R+ R, lt.R- R. lt/R/ R0 lt0R4R1 R2 llt1R# )5uG  DeepGEMM integration: fused grouped GEMM kernels from `kernels-community/deep-gemm`.

Provides:
- `deepgemm_bf16_experts_forward`: BF16 M-grouped experts forward.
- `deepgemm_fp8_fp4_linear`: end-to-end FP8/FP4 linear (BF16 in, BF16 out).
- `deepgemm_fp8_fp4_experts_forward`: FP8 (or FP4 on SM100+) M-grouped experts forward.
- `deepgemm_fp8_fp4_megamoe_experts_forward`: FP8Ã—FP4 Mega MoE forward (SM100+).

Requirements: CUDA, Hopper (SM90+), CUDA runtime â‰¥ 12.3, kernels-community/deep-gemm
â‰¥ 2.5 (Mega MoE symbols required). Mega MoE additionally needs SM100+ at call time.
)ÚannotationsN)ÚCallable)Ú	dataclass)Úlogging)Úget_cuda_runtime_versionÚis_kernels_availableÚis_torchdynamo_compilingÚresolve_internal_import)Úlazy_load_kernel)Úto_localT)Úfrozenc                  óŠ   € ] tR t^3t$ RtR]R&   R]R&   R]R&   R]R&   R]R&   R]R&   R]R	&   R]R
&   R]R&   R]R&   R]R&   RtR# )ÚDeepGEMMz>Curated entry points exposed by `kernels-community/deep-gemm`.r   Úfp8_fp4_matmulÚgrouped_fp8_fp4_matmul_ntÚgrouped_fp8_fp4_matmul_nnÚgrouped_bf16_matmul_ntÚgrouped_bf16_matmul_nnÚper_token_cast_to_fp8Ú!transform_sf_into_required_layoutÚtransform_weights_for_mega_moeÚget_symm_buffer_for_mega_moeÚfp8_fp4_mega_moeÚintÚm_alignment© N)Ú__name__Ú
__module__Ú__qualname__Ú__firstlineno__Ú__doc__Ú__annotations__Ú__static_attributes__r   ó    Ús/Volumes/fast/ai/experiments/ui-tars-smoke/.venv/lib/python3.14/site-packages/transformers/integrations/deepgemm.pyr   r   3   sI   ‡ áHàÓØ'Ó'Ø'Ó'Ø$Ó$Ø$Ó$Ø#Ó#Ø'/Ó/Ø$,Ó,Ø"*Ó*ØÓð ×r#   r   c               ó    € V ^8„  d   QhRRRR/# ©é   Úrequires_sm100ÚboolÚreturnr   r   )Úformats   "r$   Ú__annotate__r,   K   s   € ÷ Rñ R¨$ð R¸8ñ Rr#   c                óZ  € \        4       '       gç   \        4       '       g   \        R4      h\        P                  P                  4       '       g   \        R4      h\        P                  P                  4       w  rV '       d   RMRpW9  d!   V '       d   RMRp\        RV RV V R24      h\        4       w  rVV^
8X  d   R MR!pWV3V8  d,   \        RV V R	V^ ,           RV^,           R
V RV R24      h\        R4      pVf   \        R4      h\        VRR4      p	\        VRR4      p
\        VRR4      p\        VRR4      p\        VRR4      p\        VRR7      p\        VRR4      p\        VRR4      p\        VRR4      p\        VRR4      p\        VRR4      pRV	3RV
3RV3RV3RV3RV3RV3RV3RV3RV3RV33 UUu. uF  w  ppVe   K  VNK  	  pppV'       d   \        RRP                  V4       R24      h\        V	V
VVVVVVVV\        V! 4       4      R7      # u uppi )"zäLoad DeepGEMM once; raise `ImportError` if env or any required symbol is missing.

`requires_sm100` raises a Blackwell-specific error for callers (FP4 / Mega MoE)
that won't work on Hopper, instead of the generic SM90+ message.
zYDeepGEMM kernel requires the `kernels` package. Install it with `pip install -U kernels`.z9DeepGEMM kernel requires CUDA, but CUDA is not available.zBlackwell (SM100)z"Hopper (SM90) or Blackwell (SM100)zDeepGEMM requires z; current device is SMÚ.zDeepGEMM on SMu    requires CUDA runtime â‰¥ z, found z	deep-gemmNuc   Failed to load `kernels-community/deep-gemm` â€” check that a build matches the current torch/CUDA.Úfp8_fp4_gemm_ntÚ$m_grouped_fp8_fp4_gemm_nt_contiguousÚ$m_grouped_fp8_fp4_gemm_nn_contiguousÚ!m_grouped_bf16_gemm_nt_contiguousÚ!m_grouped_bf16_gemm_nn_contiguouszutils.per_token_cast_to_fp8)Úchained_pathr   r   r   Ú&get_mk_alignment_for_contiguous_layoutr   z-DeepGEMM kernel is missing required symbols: z, z'. Update with `pip install -U kernels`.)r   r   r   r   r   r   r   r   r   r   r   )é
   )é	   r6   )é   r7   )r8   é   )r   r   ÚImportErrorÚtorchÚcudaÚis_availableÚget_device_capabilityr   r
   Úgetattrr	   Újoinr   r   )r(   ÚmajorÚminorÚallowedÚarchÚ
cuda_majorÚ
cuda_minorÚmin_cudaÚkernelr   r   r   r   r   r   r   r   r   Úget_mk_alignmentr   ÚnameÚattrÚmissings   &                      r$   Ú_load_deepgemm_kernelrM   J   sÞ  € ô $×%Ò%Ü#×%Ò%ÜØkóð ô �z‰z×&Ñ&×(Ò(ÜÐYÓZÐZä—z‘z×7Ñ7Ó9‰ˆ÷ *‘%¨wˆØÔß*8Ñ&Ð>bˆDÜÐ 2°4°&Ð8NÈuÈgÐV[ÐU\Ð\]Ð^Ó_Ð_ô ":Ó!;Ñˆ
Ø# rœk‘7¨wˆØÐ# hÔ.ÜØ   ¨ wÐ.IÈ(ÐSTÍ+ÈÐVWÐX`ÐabÕXcÐWdð eØ#˜ A j \°ð4óð ô
 ˜kÓ*€FØ‚~ÜØqó
ð 	
ô ˜VÐ%6¸Ó=€NÜ '¨Ð0VÐX\Ó ]ÐÜ '¨Ð0VÐX\Ó ]ÐÜ$ VÐ-PÐRVÓWÐÜ$ VÐ-PÐRVÓWÐÜ3°FÐIfÔgÐÜ(/°Ð8[Ð]aÓ(bÐ%Ü%,¨VÐ5UÐW[Ó%\Ð"Ü#*¨6Ð3QÐSWÓ#XÐ Ü˜vÐ'OÐQUÓVÐÜ˜vÐ'9¸4Ó@Ðð
  Ð/Ø3Ð5NÐOØ3Ð5NÐOØ0Ð2HÐIØ0Ð2HÐIØ*Ð,AÐBØ0Ð2SÐTØ-Ð/MÐNØ+Ð-IÐJØ5Ð7GÐHØÐ!1Ð2ñ
ôñ
‰JˆD�$ð ÷ 	ˆñ
ð ñ ÷" ÜØ;¸D¿I¹IÀgÓ<NÐ;OÐOvÐwó
ð 	
ô Ø%Ø";Ø";Ø5Ø5Ø3Ø*KØ'EØ%AØ)ÜÑ(Ó*Ó+ôð ùó+s   Ç	H'ÇH'c               ó    € V ^8„  d   QhRRRR/# )r'   r(   r)   r*   ÚNoner   )r+   s   "r$   r,   r,   ¡   s   € ÷ ñ ¨dð ¸tñ r#   c                ó   € \        V R 7      pR# )©r(   N)rM   )r(   Ú_s   & r$   Ú_populate_deepgemm_kernelrS       s   € ä¨^Ô<€AÙr#   c               ó    € V ^8„  d   QhRRRR/# r&   r   )r+   s   "r$   r,   r,   ¦   s   € ÷ @ñ @¨ð @¸(ñ @r#   c                óR   € \        4       '       d   \        V R 7       \        V R 7      # )rQ   )r   rS   rM   rQ   s   &r$   Úload_deepgemm_kernelrV   ¦   s   € Ü×!Ò!Ü!°Õ@Ü °Ô?Ð?r#   c               ó    € V ^8„  d   QhRRRR/# )r'   Údeviceútorch.devicer*   r)   r   )r+   s   "r$   r,   r,   °   s   € ÷ =ñ =�lð = tñ =r#   c                óT   € \         P                  P                  V 4      ^ ,          ^
8¬  # )z�``True`` for Blackwell (SM100+). Cached: device capability is fixed for the
process lifetime and this gets hit on every linear/expert forward.
)r;   r<   r>   ©rX   s   &r$   Ú	_is_sm100r\   ¯   s#   € ô
 �:‰:×+Ñ+¨FÓ3°AÕ6¸"Ñ<Ð<r#   zset[int]Ú_DEEPGEMM_VISITED_DEVICESc               ó$   € V ^8„  d   QhRRRRRR/# )r'   rX   rY   ÚcontextÚstrr*   rO   r   )r+   s   "r$   r,   r,   º   s&   € ÷ !Oñ !O ,ð !O¸ð !OÀñ !Or#   c                ó"  € V P                   e   V P                   M\        P                  P                  4       p\        P                  V4       \        \        4      ^8:  d   R# RpVR8X  d   \        VR,           4      h\        VR,           4      h)u6  Reject DeepGEMM calls that span multiple CUDA devices in the same process
(e.g. ``device_map="auto"`` across N GPUs). DeepGEMM loads each kernel via
``cuKernelGetFunction``, which binds the resulting ``CUfunction`` handle to
the CUDA context that was current at load time. Driving the same cached
handle from a different device's context launches it against the wrong
module/context and produces garbage. Distributed setups (torchrun + TP/EP)
don't trip this because each process owns exactly one device's context.

The fix is a build-time choice on the DeepGEMM side: compiling with
``DG_JIT_USE_RUNTIME_API=1`` swaps the loader for the runtime API
(context-free ``cudaKernel_t``) and lifts the restriction â€” but it has to
be baked into the wheel, setting the env var at Python runtime won't change
the loader the cached ``.so`` already uses. Until the kernels-community build
we ship picks that up, we reject single-process multi-device by default.

Raised as :class:`ImportError` from the per-linear path so :func:`fp8_linear`
falls back to Triton (which loads through the runtime API and has no such
binding); raised as :class:`RuntimeError` from the experts path where there's
no fallback â€” the user explicitly chose ``experts_implementation="deepgemm"``
and must switch to ``"grouped_mm"`` / ``"eager"`` or run distributed.
NzäDeepGEMM caches each kernel's `CUfunction` against the CUDA context it was first loaded under, so driving it from a different device in the same process produces garbage. Run distributed (TP/EP) so each process owns one device, Úlinearz:or fall back to the Triton kernel (handled automatically).z.or pick `experts_implementation='grouped_mm'`.)	Úindexr;   r<   Úcurrent_devicer]   ÚaddÚlenr:   ÚRuntimeError)rX   r_   ÚidxÚmsgs   &&  r$   Ú_assert_single_devicerj   º   s{   € ð, !Ÿ,™,Ò2ˆ&�,Š,¼¿
¹
×8QÑ8QÓ8S€CÜ×!Ñ! #Ô&Ü
Ô$Ó%¨Ô*Ùð	Mð ð
 �(ÔÜ˜#Ð \Õ\Ó]Ð]Ü
�sÐMÕMÓ
NÐNr#   c               ó    € V ^8„  d   QhRRRR/# )r'   Úsfútorch.Tensorr*   r   )r+   s   "r$   r,   r,   Þ   s   € ÷ 
Yñ 
Y�|ð 
Y¨ñ 
Yr#   c                óª   € V P                  \        P                  4      pVR,           P                  R4      P                  \        P                  4      # )u­  Round each fp32 SF up to the nearest power of 2 (zero mantissa).

Mirrors `deep_gemm.utils.math.ceil_to_ue8m0`. On SM100 the kernel's
`pack_fp32_into_ue8m0` cleanly extracts the biased exponent only when the
mantissa is already zero â€” its inner shifts (`>> 15`, `>> 7`, `<< 1`)
otherwise leak mantissa bits into adjacent UE8M0 byte slots and silently
corrupt the SF. SM90 consumes raw fp32 SFs without going through this path.
iÿÿ i  €ÿ)Úviewr;   Úint32Úbitwise_and_Úfloat)rl   Úint_views   & r$   Ú_ceil_to_ue8m0rt   Þ   s<   € ð �w‰w”u—{‘{Ó#€HØ˜Õ&×4Ñ4Ð5EÓF×KÑKÌEÏKÉKÓXÐXr#   c               ó$   € V ^8„  d   QhRRRRRR/# )r'   rl   rm   Úexpected_mnz
int | Noner*   r   )r+   s   "r$   r,   r,   ë   s!   € ÷ -ñ -˜lð -¸ð -È|ñ -r#   c                óú  € \        V P                  4      pV P                  \        P                  8X  dŒ   VeA   V P                  R4      V8  d+   WP                  R4      ,          pV P                  VRR7      p V'       d/   V P                  4       P                  \        P                  4      p MCV P                  4       p M2V P                  \        P                  8X  d   V'       d   \        V 4      p V P                  4       R9  d   \        RV P                  4        R24      hV P                  R4      pV P                  R4      p^V P                  4       ,          pV) V,          ) V,          pV P                  4       ^8X  d   ^V3M
WW,          ^V3p\!        V P#                  4       4      V8X  d   V # \        P$                  ! V P&                  W€P                  V P                  R7      p	V	P)                  V 4       V	# )uT  Lay out `sf` as DeepGEMM's `check_sf_layout` expects: MN-major
(`stride(-2) == 1`) and TMA-aligned (`stride(-1) == align(mn, 16/esize)`).

Inputs come in three flavors:
  - `float8_e8m0fnu` on SM100: raw UE8M0 bytes â€” pack 4 K-bytes â†’ int32
    (last dim /4) for the kernel's `(INT, 1, gran_k)` path.
  - `float8_e8m0fnu` on SM90: SM90 dispatch only accepts FP32 SFs, so cast
    UE8M0 â†’ FP32 (exact upcast â€” UE8M0 is the biased-exponent half of a
    pow-of-2 FP32, so `.float()` rebuilds the original FP32 scale exactly).
  - `float32`: per-token / per-block SFs from `per_token_cast_to_fp8` or
    on-disk weights â€” round to UE8M0 on SM100 (see `_ceil_to_ue8m0`).
  - `int32`: already-packed UE8M0 â€” pass through.

When `expected_mn` is set and the SF's M-dim is smaller (block-quantized
UE8M0, e.g. DSv4-Flash compressor weights with `(N/128, K/128)` SFs), we
repeat the SF on the M-axis to per-row before packing â€” the `(INT, 1, gran_k)`
DeepGEMM kernel branch is the only UE8M0 path on SM100; for `gran_mn > 1`
the kernel only handles FP32 SFs and would otherwise reject our INT SF here.
©Údimz"DeepGEMM SF must be 2D or 3D, got ÚD©ÚdtyperX   éþÿÿÿ)r'   r9   éÿÿÿÿ)r\   rX   r|   r;   Úfloat8_e8m0fnuÚsizeÚrepeat_interleaveÚ
contiguousro   rp   rr   Úfloat32rt   ry   Ú
ValueErrorÚelement_sizeÚtupleÚstrideÚempty_stridedÚshapeÚcopy_)
rl   rv   Úis_sm100Úgran_mnÚmnÚkfÚalign_toÚ
aligned_mnÚtarget_stridesÚouts
   &&        r$   Ú_coerce_sf_for_kernelr“   ë   sr  € ô( ˜Ÿ™Ó#€HØ	‡x�x”5×'Ñ'Ô'ØÒ" r§w¡w¨r£{°[Ô'@Ø!§W¡W¨R£[Õ0ˆGØ×%Ñ% g°2Ð%Ó6ˆBßØ—‘“×%Ñ%¤e§k¡kÓ2‰Bà—‘“‰BØ	�‰”U—]‘]Ô	"§xÜ˜BÓˆà	‡v�vƒx�vÔÜÐ=¸b¿f¹f»h¸ZÀqÐIÓJÐJà	�‰�‹€BØ	�‰�‹€BØ�R—_‘_Ó&Õ&€HØ�3˜(•?Ð# hÕ.€JØ(*¯©«°A¬�a˜‘_¸B½OÈQÐPZÐ;[€NäˆR�Y‰Y‹[Ó˜^Ô+Øˆ	Ü
×
Ò
˜bŸh™h¨¿h¹hÈrÏyÉyÔ
Y€CØ‡I�Iˆb„MØ€Jr#   c          
     ó,   € V ^8„  d   QhRRRRRRRRRR	/# )
r'   Úweightrm   Úweight_scale_invÚ
block_sizeztuple | Noner‹   r)   r*   Údictr   )r+   s   "r$   r,   r,     s4   € ÷ /ñ /Øð/Ø,8ð/ØFRð/Ø^bð/à	ñ/r#   c                ó  € V P                   \        P                  8X  d	   RRR^ RR/# Vf   \        R4      h\	        V4      pVR	9  d   \        RV R24      hVP                   \        P
                  8X  d   V'       d	   RRR^€RR/# RRR^€/# )
u#  Pick the `per_token_cast_to_fp8` kwargs from weight dtype + SF dtype + arch.

Cases mirror the kernel's recipes:
  - FP4 weights (`int8`): gran_k=32 packed-UE8M0 SF. SM100+ only.
  - FP8 weights + UE8M0 SF on SM100: gran_k=128 packed-UE8M0 SF (DSv4).
  - FP8 weights + UE8M0 SF on SM90: gran_k=128 FP32 SF â€” the SM90 dispatch in
    `layout.hpp` only matches FP32 SFs, so we keep act SFs as FP32 (and float
    the weight SF in `_coerce_sf_for_kernel`; UE8M0 â†’ FP32 is an exact upcast).
  - FP8 weights + float SF: gran_k=128 float SF (DSv3).
Ú	use_ue8m0TÚgran_kÚuse_packed_ue8m0z]DeepGEMM requires block-wise quantized FP8 weights, but the experts have no `block_size` set.u?   DeepGEMM requires `block_size` âˆˆ {(128, 128), (1, 128)}, got r.   F))é€   r�   )é   r�   )r|   r;   Úint8r„   r†   r   )r•   r–   r—   r‹   s   &&&&r$   Ú_select_fp8_cast_kwargsr      s¤   € ð ‡|�|”u—z‘zÔ!Ø˜T 8¨RÐ1CÀTÐJÐJàÒÜØkó
ð 	
ô �zÓ"€JØÐ/Ô/ÜÐ\Ð]gÐ\hÐhiÐjÓkÐkØ×Ñ¤×!5Ñ!5Ô5¿(Ø˜T 8¨SÐ2DÀdÐKÐKØ˜ ¨#Ð.Ð.r#   c          
     ó,   € V ^8„  d   QhRRRRRRRRRR	/# )
r'   Úexpert_ids_sortedrm   Únum_expertsr   Ú	alignmentÚuse_psum_layoutr)   r*   z&tuple[torch.Tensor, torch.Tensor, int]r   )r+   s   "r$   r,   r,   :  s4   € ÷ ?ñ ?Ø#ð?Ø25ð?ØBEð?ØX\ð?à+ñ?r#   c                óø  € V P                   pV P                  ^ 4      p\        P                  ! V P	                  4       V^ V^,
          R7      P                  4       pWb,           ^,
          V,          V,          pV\        WQ4      V^,
          ,          ,           pWv,
          p	\        P                  P                  P                  V	P                  ^ 4      R4      p
\        P                  ! WTR7      W ,          ,           pV'       d!   VP                  ^ 4      P	                  4       pMS\        P                  ! V3RV\        P                  R7      p\        P                  ! W8  V P	                  4       R4      WË&   W¼V3# )av  Build the TMA-aligned grouped layout DeepGEMM expects.

Returns `(sorted_to_padded, grouped_layout, total_padded_rows)`:
  - `grouped_layout` is per-row expert id (Hopper, with `-1` for padding /
    sentinels) or a cumsum of aligned per-expert counts (Blackwell).
  - EP sentinels (values == `num_experts`) are routed past the last expert
    block so DeepGEMM skips them.
)ÚbinsÚminÚmaxr[   ©rX   r|   )rž   é    r~   )rX   r€   r;   Úhistcr   Úlongr¨   ÚnnÚ
functionalÚpadÚcumsumÚarangeÚfullrp   Úwhere)r¢   r£   r¤   r¥   rX   Ú
num_tokensÚtokens_per_expertÚaligned_tokens_per_expertÚtotal_padded_rowsÚpadding_per_expertÚcumulative_paddingÚsorted_to_paddedÚgrouped_layouts   &&&&         r$   Ú!_build_deepgemm_contiguous_layoutr½   :  s9  € ð ×%Ñ%€FØ"×'Ñ'¨Ó*€JäŸšÐ$5×$9Ñ$9Ó$;À+ÐSTÐZeÐhiÕZiÔj×oÑoÓqÐØ"3Õ"?À!Õ"CÈ	Õ!QÐU^Õ ^Ðà"¤S¨Ó%AÀYÐQRÅ]Õ%SÕSÐð 3ÕFÐÜŸ™×,Ñ,×0Ñ0Ð1C×1JÑ1JÈ1Ó1MÈvÓVÐÜ—|’| JÔ>ÐASÕAfÕfÐçØ2×9Ñ9¸!Ó<×@Ñ@ÓB‰äŸšÐ%6Ð$8¸"ÀVÔSX×S^ÑS^Ô_ˆÜ+0¯;ª;Ð7HÑ7VÐXi×XmÑXmÓXoÐqsÓ+tˆÑ(àÐ->Ð>Ð>r#   c               ó(   € V ^8„  d   QhRRRRRRRR/# )r'   Úxrm   r»   r¸   r   r*   r   )r+   s   "r$   r,   r,   \  s*   € ÷ ñ ˜ð ¸ð ÐZ]ð Ðbnñ r#   c                ó�   € \         P                  ! V.V P                  R,          O5RV P                  RV P                  / pWV&   V# )z;Pad a sorted tensor into the TMA-aligned contiguous layout.:rž   NNrX   r|   )r;   Úemptyr‰   rX   r|   )r¿   r»   r¸   Úpaddeds   &&& r$   Ú_pad_for_deepgemmrÃ   \  sA   € ä�[Š[Ð*ÐY¨Q¯W©W°R­[ÒYÀÇÁÐYÐQR×QXÑQXÑY€FØ ÐÑØ€Mr#   c               ó$   € V ^8„  d   QhRRRRRR/# )r'   Úx_paddedrm   r»   r*   r   )r+   s   "r$   r,   r,   c  s#   € ÷ &ñ &°\ð &ÐUað &Ðfrñ &r#   c                ó   € W,          # ©Nr   )rÅ   r»   s   &&r$   Ú&_unpad_from_deepgemm_contiguous_layoutrÈ   c  s   € ØÕ%Ð%r#   c               ó4   € V ^8„  d   QhRRRRRRRRRRRR	R
R/# )r'   Úhidden_statesrm   Útop_k_indexÚtop_k_weightsr£   r   r   r¥   r)   r*   r†   r   )r+   s   "r$   r,   r,   j  sN   € ÷ /ñ /Øð/àð/ð  ð/ð ð	/ð
 ð/ð ð/ð ñ/r#   c                óP  € VP                  R4      pVP                  R4      pVP                  R4      p\        P                  ! V4      w  ršW
V,          ,          pWŠ,          p\	        W“WE4      w  rÞpW“8¬  P                  R4      pV	P                  V^,
          R7       VVV	VV
VVV3# )zãSort tokens by expert id and build the M-grouped padded layout.

Returns `(sorted_hidden_states_g, sample_weights_g, expert_ids_g,
          sentinel_mask, perm, sorted_to_padded, grouped_layout,
          total_padded_rows)`.
)r©   r~   )r€   Úreshaper;   Úsortr½   Ú	unsqueezeÚclamp_)rÊ   rË   rÌ   r£   r   r¥   Ú	num_top_kÚ
expert_idsÚsample_weightsÚexpert_ids_gÚpermÚsorted_hidden_states_gÚsample_weights_gr»   r¼   r¸   Úsentinel_masks   &&&&&&           r$   Ú_dispatch_routed_inputrÚ   j  sÃ   € ð × Ñ  Ó$€IØ×$Ñ$ RÓ(€JØ"×*Ñ*¨2Ó.€Nô Ÿš JÓ/Ñ€LØ*°9Õ+<Õ=ÐØ%Õ+Ðô
 ;\Ø ;ó;Ñ7ÐÐ&7ð "Ñ0×;Ñ;¸BÓ?€MØ×Ñ˜K¨!�OÐÔ,àØØØØØØØð	ð 	r#   c               ó@   € V ^8„  d   QhRRRRRRRRRRRRR	RR
RRRRR/
# )r'   Ú
out_paddedrm   Úsorted_weightsrÙ   rÖ   r»   rµ   r   rÒ   Ú
hidden_dimÚ	out_dtypeútorch.dtyper*   r   )r+   s   "r$   r,   r,   œ  sx   € ÷ _ñ _Øð_à ð_ð  ð_ð ð	_ð
 #ð_ð ð_ð ð_ð ð_ð ð_ð ñ_r#   c	                óœ  € \        W4      p	W‘P                  V	P                  4      P                  R4      ,          p
V
P	                  VR4       \
        P                  ! V4      p\
        P                  ! VP                  ^ 4      V	P                  R7      W³&   W«,          P                  WVV4      P                  ^R7      P                  V4      # )uR   Unpad â†’ weighted multiply â†’ mask sentinels â†’ restore order â†’ top-k reduce.g        r[   rx   r~   )rÈ   Útor|   rÐ   Úmasked_fill_r;   Ú
empty_liker²   r€   rX   ro   Úsum)rÜ   rÝ   rÙ   rÖ   r»   rµ   rÒ   rÞ   rß   r’   ÚweightedÚinv_perms   &&&&&&&&&   r$   Ú_combine_routed_outputrè   œ  s    € ô 1°Ó
N€CØ×&Ñ& s§y¡yÓ1×;Ñ;¸BÓ?Õ?€Hð ×Ñ˜-¨Ô-Ü×Ò Ó%€HÜ—\’\ $§)¡)¨A£,°s·z±zÔB€H�NàÕ×"Ñ" :¸*ÓE×IÑIÈaÐIÓP×SÑSÐT]Ó^Ð^r#   c               ó8   € V ^8„  d   QhRRRRRRRRRRR	R
RRRR/# )r'   Úinputrm   r•   r–   Úbiasztorch.Tensor | Noner—   ztuple[int, int] | NoneÚoutput_dtyperà   Úactivation_scaler*   r   )r+   s   "r$   r,   r,   ¶  sX   € ÷ (ñ (Øð(àð(ð #ð(ð ð	(ð
 'ð(ð ð(ð *ð(ð ñ(r#   c           
     óâ  € \        V P                  RR7       Ve   \        R4      hV P                  \        P
                  \        P                  39  d   \        RV P                   24      h\        VP                  \        P                  8H  R7      p\        WV\        V P                  4      4      pV P                  RV P                  R,          4      p	VP                  ! V	3/ VB w  r«\        P                  ! V
P                  ^ ,          VP                  ^ ,          V P                  VR7      pVP!                  R4      '       d   ^^VR	,          3MRpVP#                  V
\%        WºP'                  ^ 4      R
7      3V\%        W!P'                  ^ 4      R
7      3VVR7       VP                  V P                  RR VP                  ^ ,          3,           4      pVe   VP)                  V4       V# )uç   End-to-end DeepGEMM linear: per-token activation quant + FP8/FP4 matmul.

Static (per-tensor) activation quantization is rejected â€” DeepGEMM needs
per-row SFs. Callers should route static activations through the Triton fallback.
rb   ©r_   Nz@DeepGEMM linear does not support static activation quantization.z7DeepGEMM linear requires FP16 or BF16 activations, got rQ   rª   rœ   r›   ©rv   )Úreciper~   )rj   rX   ÚNotImplementedErrorr|   r;   Úbfloat16Úfloat16r„   rV   rŸ   r    r\   ro   r‰   r   rÁ   Úgetr   r“   r€   Úadd_)rê   r•   r–   rë   r—   rì   rí   ÚdeepgemmÚcast_kwargsÚinput_2dÚ	qinput_2dÚscale_2dÚoutputÚ	sf_recipes   &&&&&&&       r$   Údeepgemm_fp8_fp4_linearrþ   ¶  s“  € ô ˜%Ÿ,™,°Õ9àÒ#Ü!Ð"dÓeÐeØ‡{�{œ5Ÿ>™>¬5¯=©=Ð9Ô9ÜÐRÐSX×S^ÑS^ÐR_Ð`ÓaÐaä#°6·<±<Ä5Ç:Á:Ñ3MÔN€HÜ)¨&ÀJÔPYÐZ_×ZfÑZfÓPgÓh€Kà�z‰z˜"˜eŸk™k¨"�oÓ.€HØ"×8Ò8¸ÑQÀ[ÑQÑ€IÜ�[Š[˜Ÿ™¨Õ+¨V¯\©\¸!­_ÀUÇ\Á\ÐYeÔf€Fð 2=·±ÐAS×1TÒ1T��A�{ 8Õ,Ñ-ÐZ^€IØ×ÑØ	Ô)¨(ÇÁÈqÓ@QÔRÐSØ	Ô&Ð'7Ç[Á[ÐQRÃ^ÔTÐUØØð	 ô ð �[‰[˜Ÿ™ S bÐ)¨V¯\©\¸!­_Ð,>Õ>Ó?€FØÒØ�‰�DÔØ€Mr#   c          
     ó,   € V ^8„  d   QhRRRRRRRRRR/# ©r'   Úselfútorch.nn.ModulerÊ   rm   rË   rÌ   r*   r   )r+   s   "r$   r,   r,   á  s:   € ÷ >ñ >Ø
ð>àð>ð ð>ð  ð	>ð
 ñ>r#   c                óì  € VP                   \        P                  8w  d   \        R VP                    24      h\	        4       pV P
                  '       d   VP                  MVP                  pVP                  pVP                  R4      pVP                  ^ 4      pVP                  R4      p	\        WW0P                  VP                  \        V4      4      w  p
ppppppp\        V P                  '       d   V P                   MV P"                  4      p\        V P$                  4      pV P&                  '       d4   \        V P                  '       d   V P(                  MV P*                  4      MRpV P&                  '       d   \        V P,                  4      MRpV P
                  '       d   VP.                  R,          MVP.                  ^,          p\1        W¯V4      p\        P2                  ! VVWaP                   R7      pV! VVVV\        V4      R7       V P&                  '       d   VP5                  ^ VVV,          4       V P                  '       d   V P7                  V4      MV P9                  V4      p\        P2                  ! VW–VP                   R7      pV! VVVV\        V4      R7       V P&                  '       d   VP5                  ^ VVV,          4       \;        VVVVVVVV	VP                   4	      # )ú;DeepGEMM experts path requires bfloat16 hidden states, got Nrª   )r¥   r~   )r|   r;   ró   r„   rV   Úis_transposedr   r   rX   r€   rÚ   r£   r   r\   r   Úhas_gateÚgate_up_projÚup_projÚ	down_projÚhas_biasÚgate_up_proj_biasÚup_proj_biasÚdown_proj_biasr‰   rÃ   rÁ   Ú
index_add_Ú_apply_gateÚact_fnrè   )r  rÊ   rË   rÌ   r÷   Úgrouped_bf16_matmulrX   rÒ   rµ   rÞ   Úsorted_hiddenrÝ   rÕ   rÙ   rÖ   r»   r¼   r¸   Ú	weight_upÚweight_downÚup_biasÚ	down_biasÚ
up_out_dimÚactÚproj_outr’   s   &&&&                      r$   Údeepgemm_bf16_experts_forwardr  á  st  € ð ×ÑœeŸn™nÔ,ÜÐVÐWd×WjÑWjÐVkÐlÓmÐmä#Ó%€Hà=A×=O×=OÐ=O˜(×9Ò9ÐU]×UtÑUtÐà×!Ñ!€FØ× Ñ  Ó$€IØ×#Ñ# AÓ&€JØ×#Ñ# BÓ'€Jô 	Ø M×3CÑ3CÀX×EYÑEYÔ[dÐekÓ[ló	ñ	ØØØØØØØØô
 ¨d¯m¯m¨m˜×*Ò*ÀÇÁÓN€IÜ˜4Ÿ>™>Ó*€KØZ^×Zg×ZgÐZgŒh°··°�t×-Ò-ÀD×DUÑDUÔVÐmq€GØ15··°”˜×,Ñ,Ô-ÀD€Ið )-×(:×(:Ð(:�—‘ Ö$À	ÇÁÐPQÕ@R€JÜ
˜MÐ=NÓ
O€CÜ�{Š{Ð,¨jÀ×ObÑObÔc€HÙ˜˜Y¨°.ÔR[Ð\bÓRcÕdØ‡}‡}€}Ø×Ñ˜AÐ/°¸Õ1FÔGà-1¯]¯]¨]ˆt×Ñ Ô)ÀÇÁÈHÓ@U€Hô �+Š+Ð'¨È-×J]ÑJ]Ô
^€CÙ˜ +¨s°NÔT]Ð^dÓTeÕfØ‡}‡}€}Ø�‰�qÐ*¨I°lÕ,CÔDä!ØØØØØØØØØ×Ñó
ð 
r#   c          
     ó,   € V ^8„  d   QhRRRRRRRRRR/# r   r   )r+   s   "r$   r,   r,   "  sA   € ÷ Pñ PØ
ðPàðPð ðPð  ð	Pð
 ñPr#   c                ó  € \        VP                  R R7       V P                  R8X  d   \        R4      hVP                  \
        P                  8w  d   \        RVP                   24      h\        V P                  P                  \
        P                  8H  R7      pV P                  '       d   VP                  MVP                  pVP                  pVP                  R4      pVP                  ^ 4      pVP                  R4      p	\        V P                   '       d   V P"                  MV P$                  4      p
\        V P                   '       d   V P&                  MV P(                  4      p\        V P                  4      p\        V P*                  4      p\-        W«V P.                  \1        V4      4      p\3        WW0P4                  VP6                  \1        V4      4      w  ppppppppVP9                  R4      '       d   ^^VR,          3MRpVP:                  ! V3/ VB w  pp\=        VVV4      p\=        VVV4      p\
        P>                  ! VV
P@                  ^,          V\
        P                  R	7      pV! V\C        VVR
7      3V
\C        WºP                  R4      R
7      3VVV\1        V4      R7       V P                   '       d   V PE                  V4      MV PG                  V4      pVP:                  ! V3/ VB w  pp\
        P>                  ! VW–\
        P                  R	7      pV! V\C        VVR
7      3V\C        WÜP                  R4      R
7      3VVV\1        V4      R7       \I        VVVVVVVV	VP                  4	      # )Úexpertsrï   ÚstaticzJDeepGEMM experts dispatch does not support static activation quantization.r  rQ   rœ   r›   Nrª   rð   )rñ   r¥   r~   r}   )%rj   rX   Úactivation_schemerò   r|   r;   ró   r„   rV   r	  rŸ   r  r   r   r€   r   r  r  r  Úgate_up_proj_scale_invÚup_proj_scale_invÚdown_proj_scale_invr    r—   r\   rÚ   r£   r   rõ   r   rÃ   rÁ   r‰   r“   r  r  rè   )r  rÊ   rË   rÌ   r÷   Úgrouped_fp8_fp4_matmulrX   rÒ   rµ   rÞ   r  Úweight_scale_upr  Úweight_scale_downrø   r  rÝ   Ú_expert_ids_grÙ   rÖ   r»   r¼   r¸   rý   Úact_fp8Ú
act_scalesr  Úproj_fp8Úproj_scalesr’   s   &&&&                          r$   Ú deepgemm_fp8_fp4_experts_forwardr+  "  s%  € ô ˜-×.Ñ.¸	ÕBà×Ñ Ô)Ü!Ð"nÓoÐoØ×ÑœeŸn™nÔ,ÜÐVÐWd×WjÑWjÐVkÐlÓmÐmä#°4·>±>×3GÑ3GÌ5Ï:É:Ñ3UÔV€Hà.2×.@×.@Ð.@ˆ×*Ò*Àh×FhÑFhð ð ×!Ñ!€FØ× Ñ  Ó$€IØ×#Ñ# AÓ&€JØ×#Ñ# BÓ'€Jä¨d¯m¯m¨m˜×*Ò*ÀÇÁÓN€IÜ¸d¿m¿m¸m˜t×:Ò:ÐQU×QgÑQgÓh€OÜ˜4Ÿ>™>Ó*€KÜ  ×!9Ñ!9Ó:Ðä)¨)ÀdÇoÁoÔW`ÐagÓWhÓi€Kô 	Ø M×3CÑ3CÀX×EYÑEYÔ[dÐekÓ[ló	ñ	ØØØØØØØØð 2=·±ÐAS×1TÒ1T��A�{ 8Õ,Ñ-ÐZ^€Ið #×8Ò8¸ÑVÈ+ÑVÑ€GˆZÜ Ð)9Ð;LÓM€GÜ" :Ð/?ÐARÓS€JÜ�{Š{Ð,¨i¯o©o¸aÕ.@ÈÔW\×WeÑWeÔf€HÙØ	Ô'¨
Ð@QÔRÐSØ	Ô)¨/Ç~Á~ÐVXÓGYÔZÐ[ØØØÜ! &Ó)õð .2¯]¯]¨]ˆt×Ñ Ô)ÀÇÁÈHÓ@U€Hð %×:Ò:¸8ÑSÀ{ÑSÑ€HˆkÜ
�+Š+Ð'¨Ì%Ï.É.Ô
Y€CÙØ	Ô(¨ÐBSÔTÐUØ	Ô+Ð,=×K[ÑK[Ð\^ÓK_Ô`ÐaØØØÜ! &Ó)õô "ØØØØØØØØØ×Ñó
ð 
r#   c               ó    € V ^8„  d   QhRRRR/# )r'   Úmoduler  r*   rO   r   )r+   s   "r$   r,   r,   u  s   € ÷ 7Rñ 7R /ð 7R°dñ 7Rr#   c                ój  € \        RR7      p\        V P                  P                  4      p\        V P                  P                  4      p\        V P
                  P                  4      P                  \        P                  4      P                  4       p\        V P                  P                  4      P                  \        P                  4      P                  4       pV P                  pV P                  pV P                  pV^ ,          ^ 8w  g   V^ ,          ^ 8w  d   \        RV RV R24      hVP                  VP!                  4       ^V,          VR
VR7      p	VP                  VP!                  4       VVR
VR7      p
VP#                  WI3WZ34      w  w  r¹w  rÊ\        P$                  P'                  VRR7      V n        \        P$                  P'                  V	RR7      V n        \        P$                  P'                  VRR7      V n
        \        P$                  P'                  V
RR7      V n        R	# )uí  One-shot pack + permute of an FP8Experts module's L1/L2 weights into the
Mega MoE UTCCP layout. Called lazily on the first megamoe forward; idempotent
via the caller's ``_megamoe_transformed`` flag.

Steps:
  1. Cast UE8M0 SF â†’ FP32 and call ``transform_sf_into_required_layout`` â†’
     packed int32 in MN-major TMA-aligned layout.
  2. Run ``transform_weights_for_mega_moe``: interleaves gate/up on L1 and
     transposes both SFs for UTCCP.
  3. Overwrite the loader-side parameters in place; the interleave preserves
     the ``[E_local, 2*I, *]`` leading dims so downstream ``.size(...)`` reads
     stay valid.

Unwraps any ``DTensor`` wrappers FSDP2/EP may have placed around the loader-
side Parameters â€” the kernel takes raw pointers.
TrQ   zwDeepGEMM Mega MoE requires `hidden_dim` and `intermediate_hidden` divisible by 32 (FP8 SF granularity); got hidden_dim=z, intermediate_hidden=r.   )rñ   Ú
num_groupsF)Úrequires_gradN)rž   é    )rV   r   r   Údatar"  r  ro   r;   rŸ   r‚   r	  Úintermediate_dimr£   rÞ   r„   r   rr   r   r®   Ú	Parameter)r-  r÷   Úgate_up_sf_rawÚdown_sf_rawÚ	gate_up_wÚdown_wÚintermediate_hiddenÚnum_local_expertsrÞ   Ú
gate_up_sfÚdown_sfÚgate_upÚdowns   &            r$   Úsetup_megamoe_weightsr?  u  sú  € ô" $°4Ô8€HÜ˜f×;Ñ;×@Ñ@ÓA€NÜ˜6×5Ñ5×:Ñ:Ó;€Kä˜×,Ñ,×1Ñ1Ó2×7Ñ7¼¿
¹
ÓC×NÑNÓP€IÜ�f×&Ñ&×+Ñ+Ó,×1Ñ1´%·*±*Ó=×HÑHÓJ€Fà ×1Ñ1ÐØ×*Ñ*ÐØ×"Ñ"€Jà�B…˜!ÔÐ2°RÕ7¸1Ô<Üð4Ø4>°<Ð?UÐViÐUjÐjkðmó
ð 	
ð
 ×;Ñ;Ø×ÑÓØ	ÐÕØØØ$ð <ó €Jð ×8Ñ8Ø×ÑÓØØØØ$ð 9ó €Gð .6×-TÑ-TØ	ÐØ	Ðó.Ñ*Ñ€W™?˜Dô  Ÿ(™(×,Ñ,¨WÀEÐ,ÓJ€FÔÜ$)§H¡H×$6Ñ$6°zÐQVÐ$6Ó$W€FÔ!Ü—x‘x×)Ñ)¨$¸eÐ)ÓD€FÔÜ!&§¡×!3Ñ!3°GÈ5Ð!3Ó!Q€FÖr#   c               ó0   € V ^8„  d   QhRRRRRRRRRRR	R/# )
r'   r  r  rÊ   rm   rË   rÌ   Úprocess_groupz%torch.distributed.ProcessGroup | Noner*   r   )r+   s   "r$   r,   r,   ¯  sL   € ÷ Q%ñ Q%Ø
ðQ%àðQ%ð ðQ%ð  ð	Q%ð
 9ðQ%ð ñQ%r#   c                ó~  € V P                   P                  \        P                  8w  d$   \	        RV P                   P                   R24      hVf   \        R4      h\        RR7      p\        V RR4      '       g   \        V 4       RV n	        VP                  R4      pVP                  ^ 4      pVP                  R4      pV P                   P                  ^ 4      p	V P                   P                  ^4      ^,          p
W”P                  4       ,          p\        V R	R4      e   V P                  P                  V8  d   VP                  VVVVVV
R
7      V n        VP                  VR^ RR7      w  rÍV P                  P                  RV P!                  V4       V P                  P"                  RV P!                  V4       V P                  P$                  RV P!                  V4       V P                  P&                  RV P!                  V4       \        P(                  ! Wx3\        P*                  VP,                  R7      pVP/                  VV P                   V P0                  3V P2                  V P4                  3V P                  \        \        V RR4      RR4      R7       VP7                  VP                  4      # )u·  FP8 acts Ã— FP4 weights Mega MoE forward (SM100+).

Fuses EP dispatch + L1 + SwiGLU + L2 + EP combine into one kernel,
overlapping NVLink with tensor-core compute. The kernel handles the full
`(num_tokens, hidden) â†’ (num_tokens, hidden)` MoE forward including the
weighted top-k reduction; the caller must NOT all-reduce the output.

`process_group` is supplied automatically by `MoeTensorParalellExperts._prepare_input_fn`
when the module is wrapped for TP â€” it's required for the symm-buffer rendezvous
on first forward. `top_k_index` is GLOBAL expert ids (`-1` marks skipped slots).

Caller-managed `self` attributes:
  - `gate_up_proj`, `gate_up_proj_scale_inv`: L1 weight + UE8M0 SF.
  - `down_proj`, `down_proj_scale_inv`: L2 weight + UE8M0 SF.
  Both pairs must be transformed together via
  `transform_weights_for_mega_moe((gate_up, gate_up_sf), (down, down_sf))`.
  - `config.swiglu_limit` (optional): SwiGLU clamp; absent â†’ unclamped.
zJDeepGEMM Mega MoE requires FP4-packed expert weights (dtype=`int8`), got `z/`. Use the 'deepgemm' dispatch for FP8 experts.Nz©DeepGEMM Mega MoE requires a `process_group` for the EP group. The TP wrapping (MoeTensorParalellMegaMoeExperts) supplies it automatically; pass it explicitly otherwise.TrQ   Ú_megamoe_transformedFÚsymm_buffer)ÚhiddenÚnum_topkr£   Únum_max_tokens_per_rankr9  )rš   r›   rœ   r{   ÚconfigÚswiglu_limit)Úactivation_clampr~   )r  r|   r;   rŸ   rg   r„   rV   r?   r?  rC  r€   rD  rG  r   r   r¿   rŠ   Úx_sfÚtopk_idxÚtopk_weightsrÁ   ró   rX   r   r   r	  r"  râ   )r  rÊ   rË   rÌ   rA  r÷   rÒ   rµ   rÞ   r:  r9  Únum_global_expertsÚx_fp8rK  Úys   &&&&&          r$   Ú(deepgemm_fp8_fp4_megamoe_experts_forwardrQ  ¯  s�  € ð2 ×Ñ×Ñ¤%§*¡*Ô,ÜðØ×!Ñ!×'Ñ'Ð(Ð(WðYó
ð 	
ð
 ÒÜðió
ð 	
ô
 $°4Ô8€Hô �4Ð/°×7Ò7Ü˜dÔ#Ø$(ˆÔ!à× Ñ  Ó$€IØ×#Ñ# AÓ&€JØ×#Ñ# BÓ'€JØ×)Ñ)×.Ñ.¨qÓ1ÐØ×+Ñ+×0Ñ0°Ó3°qÕ8ÐØ*×-?Ñ-?Ó-AÕAÐô ˆt�] DÓ)Ò1°T×5EÑ5E×5]Ñ5]Ð`jÔ5jØ#×@Ñ@ØØØØ*Ø$.Ø 3ð Aó 
ˆÔð ×0Ñ0°È$ÐWYÐlpÐ0Óq�K€EØ×Ñ×Ñ�{˜
Ð#×)Ñ)¨%Ô0Ø×Ñ×Ñ˜+˜:Ð&×,Ñ,¨TÔ2Ø×Ñ×Ñ˜k˜zÐ*×0Ñ0°Ô=Ø×Ñ×!Ñ! + :Ð.×4Ñ4°]ÔCô 	�Š�ZÐ,´E·N±NÈ=×K_ÑK_Ô`€AØ×ÑØ	Ø	×	Ñ	˜D×7Ñ7Ð8Ø	�‰˜×1Ñ1Ð2Ø×ÑÜ ¤¨¨x¸Ó!>ÀÐPTÓUð ô ð �4‰4�×#Ñ#Ó$Ð$r#   )FrÇ   )2Ú__conditional_annotations__r    Ú
__future__r   Ú	functoolsÚcollections.abcr   Údataclassesr   r;   Úutilsr   Úutils.import_utilsr   r   r   r	   Úhub_kernelsr
   Útensor_parallelr   Ú
get_loggerr   Úloggerr   ÚcacherM   Ú_dynamoÚallow_in_graphrS   rV   r\   Úsetr]   r!   rj   rt   r“   r    r½   rÃ   rÈ   rÚ   rè   ró   rþ   r  r+  r?  rQ  )rR  s   @r$   Ú<module>ra     s3  øðô
õ #ã Ý $Ý !ã å ÷ó õ *Ý %ð 
×	Ò	˜HÓ	%€ñ
 �$Ô÷ð ó ðð, ‡�öRó ðRðj ‡�×Ñöó ð÷
@ð ‡�ô=ó ð=ñ '*£eÐ ˜8Ó +õ!OõH
Y÷-õ`/õ>?õDõ&õ/õd_ð< !%Ø)-Ø %§¡Ø,0÷(õV>õBPõf7R÷tQ%ñ Q%r#   