§
    ‚Štj­¢  ã                  óp  — d Z ddlmZ ddlZddlZddlZddlZddlZddlm	Z	 ddl
mZ ddlZddlmZ ddlmZ dd	lmZmZmZmZmZ d
dlmZ d
dlmZ  ej        e¦  «        Z ed¬¦  «         G d„ d¦  «        ¦   «         Zej        d_d„¦   «         Z ej        d`d„¦   «         Z!ej        dadbd„¦   «         Z"ej#        j$        dadcd„¦   «         Z%daddd„Z&ej        ded „¦   «         Z'dfd#„Z(dgd%„Z)dhdid(„Z*djd/„Z+dkd6„Z,dld:„Z-dmd<„Z.dndB„Z/dodL„Z0 edMdN¬O¦  «        	 	 	 	 dpdqdU„¦   «         Z1drdX„Z2drdY„Z3dsd[„Z4	 dhdtd^„Z5dS )uuS  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 (output dtype follows the input).
- `deepgemm_fp8_fp4_experts_forward`: FP8 (or FP4 on SM100+) M-grouped experts forward.
- `deepgemm_fp8_fp4_megamoe_experts_forward`: FP8xFP4 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)Údeprecate_kwarg)ÚKERNELS_MAX_VERSIONÚKERNELS_MIN_VERSIONÚis_kernels_availableÚis_torchdynamo_compilingÚresolve_internal_importé   )Úlazy_load_kernel)Úto_localT)Úfrozenc                  ó‚   — e Zd ZU dZded<   ded<   ded<   ded<   ded<   ded<   ded	<   ded
<   ded<   ded<   ded<   dS )Ú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_alignmentN)Ú__name__Ú
__module__Ú__qualname__Ú__doc__Ú__annotations__© ó    ú`/var/www/html/CA-Chatbot/venv/lib/python3.11/site-packages/transformers/integrations/deepgemm.pyr   r   9   sž   € € € € € € àHÐHàÐÐÑØ'Ð'Ð'Ñ'Ø'Ð'Ð'Ñ'Ø$Ð$Ð$Ñ$Ø$Ð$Ð$Ñ$Ø#Ð#Ð#Ñ#Ø/Ð/Ð/Ñ/Ø,Ð,Ð,Ñ,Ø*Ð*Ð*Ñ*ØÐÐÑð ÐÐÑÐÐr&   r   Úreturnú
str | Nonec                 óp  — t           j                             d¦  «        pt           j                             d¦  «        } | r| S t          j        d¦  «        }|r<t           j                             t           j                             |¦  «        ¦  «        S t           j                             d¦  «        rdS dS )u‚  Resolve the CUDA toolkit root the way DeepGEMM's JIT does:
    ``CUDA_HOME`` â†’ ``CUDA_PATH`` â†’ dir of ``which nvcc`` â†’ ``/usr/local/cuda`` (``None`` if none found).

    Mirrors DeepGEMM's own ``_find_cuda_home`` so we agree on the path it will actually use, rather than
    reusing ``torch.utils.cpp_extension.CUDA_HOME`` whose resolution inits a CUDA context (fork-unsafe).
    Ú	CUDA_HOMEÚ	CUDA_PATHÚnvccz/usr/local/cudaN)ÚosÚenvironÚgetÚshutilÚwhichÚpathÚdirnameÚisdir)Ú	cuda_homer-   s     r'   Ú_get_cuda_homer7   P   s“   € õ ”
—’˜{Ñ+Ô+ÐJ­r¬z¯~ª~¸kÑ/JÔ/J€IØð ØÐÝŒ<˜ÑÔ€DØð 6ÝŒw�Š�rœwŸš¨tÑ4Ô4Ñ5Ô5Ð5Ý	„w‡}‚}Ð&Ñ'Ô'ð !Ø Ð Øˆ4r&   útuple[int, int] | Nonec                 óÚ  — t          ¦   «         } | €dS t          j                             | d¦  «        }t          j                             |¦  «        rÕ	 t          |¦  «        5 }t          j        |¦  «        }ddd¦  «         n# 1 swxY w Y   |                     d|                     di ¦  «        ¦  «                             dd¦  «        }| 	                    d¦  «        dd…         \  }}t          |¦  «        t          |¦  «        fS # t          t          t          f$ r Y nw xY wt          j                             | d	¦  «        }t          j                             |¦  «        r­	 t          |¦  «        5 }t          j        d
|                     ¦   «         ¦  «        }ddd¦  «         n# 1 swxY w Y   |rDt          |                     d¦  «        ¦  «        t          |                     d¦  «        ¦  «        fS n# t          t          f$ r Y nw xY wt          j                             | dd¦  «        }	t          j                             |	¦  «        r˜	 t          |	¦  «        5 }t          j        d|                     ¦   «         ¦  «        }ddd¦  «         n# 1 swxY w Y   |r/t          |                     d¦  «        ¦  «        }
|
dz  |
dz  dz  fS n# t          t          f$ r Y nw xY wdS )a‹  Version of the CUDA toolkit nvcc will use, as ``(major, minor)``, read off disk without a
    subprocess from (in order) ``{CUDA_HOME}/version.json``, ``version.txt``, or the ``CUDA_VERSION``
    define in ``include/cuda.h``. ``None`` if unreadable. This is the compiler that builds the kernels,
    unlike ``torch.version.cuda`` (torch's bundled runtime, which never drives a JIT compile).
    Nzversion.jsonÚ	cuda_nvccÚcudaÚversionÚ ú.r   zversion.txtzCUDA Version (\d+)\.(\d+)r   Úincludezcuda.hz#define CUDA_VERSION (\d+)iè  é
   )r7   r.   r3   ÚjoinÚisfileÚopenÚjsonÚloadr0   Úsplitr   ÚOSErrorÚ
ValueErrorÚAttributeErrorÚreÚsearchÚreadÚgroup)r6   Úversion_jsonÚfÚ
componentsr<   ÚmajorÚminorÚversion_txtÚmatchÚcuda_hÚcuda_versions              r'   Ú_get_nvcc_versionrW   c   s�  € õ Ñ Ô €IØÐØˆtå”7—<’< 	¨>Ñ:Ô:€LÝ	„w‡~‚~�lÑ#Ô#ð ð	Ý�lÑ#Ô#ð * qÝ!œY q™\œ\�
ð*ð *ð *ñ *ô *ð *ð *ð *ð *ð *ð *øøøð *ð *ð *ð *à —n’n [°*·.².ÀÈÑ2LÔ2LÑMÔM×QÒQÐR[Ð]_Ñ`Ô`ˆGØ"Ÿ=š=¨Ñ-Ô-¨b¨q¨bÔ1‰LˆE�5Ý�u‘:”:�s 5™zœzÐ)Ð)øÝ�¥^Ð4ð 	ð 	ð 	ØˆDð	øøøõ ”'—,’,˜y¨-Ñ8Ô8€KÝ	„w‡~‚~�kÑ"Ô"ð ð	Ý�kÑ"Ô"ð J aÝœ	Ð">ÀÇÂÁÄÑIÔI�ðJð Jð Jñ Jô Jð Jð Jð Jð Jð Jð Jøøøð Jð Jð Jð Jàð @Ý˜5Ÿ;š; q™>œ>Ñ*Ô*­C°·²¸A±´Ñ,?Ô,?Ð?Ð?ð@øå�Ð$ð 	ð 	ð 	ØˆDð	øøøõ ŒW�\Š\˜) Y°Ñ9Ô9€FÝ	„w‡~‚~�fÑÔð ð	Ý�f‘”ð K Ýœ	Ð"?ÀÇÂÁÄÑJÔJ�ðKð Kð Kñ Kô Kð Kð Kð Kð Kð Kð Køøøð Kð Kð Kð Kàð IÝ" 5§;¢;¨q¡>¤>Ñ2Ô2�Ø# tÑ+¨l¸TÑ.AÀbÑ-HÐHÐHðIøõ �Ð$ð 	ð 	ð 	ØˆDð	øøøð ˆ4s¤   ÁD
 Á"BÁ7D
 ÂBÂD
 Â
BÂA>D
 Ä
D$Ä#D$Å'G< Å6(F*ÆG< Æ*F.Æ.G< Æ1F.Æ2AG< Ç<HÈHÉK É#(JÊK ÊJÊK ÊJÊ3K ËK(Ë'K(FÚrequires_sm100ÚboolúDeepGEMM | strc                ó²  — t          ¦   «         �syt          ¦   «         sdt          › dt          › dt          › d�S t          j                             ¦   «         sdS t          j                             ¦   «         \  }}| rdnd}||vr| rdnd	}d
|› d|› |› d�S |dk    rdnd}t          ¦   «         }|€d|d         › d|d         › d�S t          j
                             t          j
                             |dd¦  «        ¦  «        sd|› d|d         › d|d         › d�S t          ¦   «         }|€d|› d|d         › d|d         › d�S ||k     rAd|› |› d|d         › d|d         › d|d         › d|d         › d |› d!|d         › d|d         › d�S t          d"¦  «        }|€d#S t          |d$d¦  «        }	t          |d%d¦  «        }
t          |d&d¦  «        }t          |d'd¦  «        }t          |d(d¦  «        }t!          |d)¬*¦  «        }t          |d+d¦  «        }t          |d,d¦  «        }t          |d-d¦  «        }t          |d.d¦  «        }t          |d/d¦  «        }d0„ d$|	fd%|
fd&|fd'|fd(|fd)|fd+|fd,|fd-|fd.|fd/|ffD ¦   «         }|r1d1d2                     |¦  «        › d3t          › dt          › dt          › d�	S t#          |	|
|||||||| |¦   «         ¬4¦  «        S )5aÕ  Load DeepGEMM once or returns an error message if env or any required symbol is missing. This is wrapped in a
    function that will raise an `ImportError` with the error message. The reason we raise in the wrapper rather than
    here is that @functools.cache will only cache a return value, not an exception.

    `requires_sm100` raises a Blackwell-specific error for callers (FP4 / Mega MoE) that won't work on Hopper, instead
    of the generic SM90+ message.
    zUDeepGEMM kernel requires the `kernels` package. Please install a compatible version (z <= version < z), e.g. `pip install kernels==ú`z9DeepGEMM kernel requires CUDA, but CUDA is not available.)r@   )é	   r@   zBlackwell (SM100)z"Hopper (SM90) or Blackwell (SM100)zDeepGEMM requires z; current device is SMr>   r@   )é   r]   )r^   é   Nu(   DeepGEMM's JIT needs a CUDA toolkit â‰¥ r   r   z8, but none was found. Set `CUDA_HOME` to a CUDA toolkit.Úbinr-   z:DeepGEMM's JIT compiles with nvcc, but none was found in `u,   /bin`. Point `CUDA_HOME` at a full CUDA â‰¥ z& toolkit (not a runtime-only install).zDeepGEMM found nvcc in `u–   /bin` but could not read its CUDA version (no parseable `version.json`, `version.txt`, or `include/cuda.h`). Point `CUDA_HOME` at a complete CUDA â‰¥ z	 toolkit.zDeepGEMM on SMu    needs a CUDA â‰¥ z toolkit, but nvcc z in `u.   ` is too old. Point `CUDA_HOME` at a CUDA â‰¥ z	deep-gemmuc   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   c                ó   — g | ]	\  }}|­|‘Œ
S ©Nr%   )Ú.0ÚnameÚattrs      r'   ú
<listcomp>z)_load_deepgemm_kernel.<locals>.<listcomp>ß   s-   € ð ð ð áˆD�$ð ˆ<ð 	ð ˆ<ˆ<r&   z-DeepGEMM kernel is missing required symbols: z, z'. Please install a compatible version ()r   r   r   r   r   r   r   r   r   r   r   )r   r   r
   r	   Útorchr;   Úis_availableÚget_device_capabilityr7   r.   r3   rB   rA   rW   r   Úgetattrr   r   )rX   rQ   rR   ÚallowedÚarchÚmin_cudar6   Únvcc_versionÚkernelr   r   r   r   r   r   r   r   r   Úget_mk_alignmentr   Úmissings                        r'   Ú_load_deepgemm_kernelry   ’   sÁ  € õ $Ñ%Ô%ñ 2Ý#Ñ%Ô%ð 	ð*Ý&ð*ð *Ý6Ið*ð *å&ð*ð *ð *ðõ
 Œz×&Ò&Ñ(Ô(ð 	OØNÐNå”z×7Ò7Ñ9Ô9‰ˆˆuð *Ð6�%�%¨wˆØ˜ÐÐØ*8ÐbÐ&Ð&Ð>bˆDØS¨ÐSÐSÀEÐSÈ5ÐSÐSÐSÐSð $ ršk˜k�7�7¨wˆÝ"Ñ$Ô$ˆ	ØÐð5¸8ÀA¼;ð 5ð 5ÈÐRSÌð 5ð 5ð 5ðõ Œw�~Š~�bœgŸlšl¨9°e¸VÑDÔDÑEÔEð 	ðtÈYð tð tØ2:¸1´+ðtð tØ@HÈÄðtð tð tðõ )Ñ*Ô*ˆØÐðJ¨9ð Jð Jà%-¨a¤[ðJð Jà3;¸A´;ðJð Jð Jðð
 ˜(Ò"Ð"ðA ð A¨ð Að AÀÈ!Äð Að AÈxÐXYÌ{ð Að AØ ”?ðAð AØ%1°!¤_ðAð AØ;DðAð Aà$ QœKðAð Aà*2°1¬+ðAð Að Aðõ ˜kÑ*Ô*€FØ€~ØtÐtå˜VÐ%6¸Ñ=Ô=€NÝ '¨Ð0VÐX\Ñ ]Ô ]ÐÝ '¨Ð0VÐX\Ñ ]Ô ]ÐÝ$ VÐ-PÐRVÑWÔWÐÝ$ VÐ-PÐRVÑWÔWÐÝ3°FÐIfÐgÑgÔgÐÝ(/°Ð8[Ð]aÑ(bÔ(bÐ%Ý%,¨VÐ5UÐW[Ñ%\Ô%\Ð"Ý#*¨6Ð3QÐSWÑ#XÔ#XÐ Ý˜vÐ'OÐQUÑVÔ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ð
ðñ ô €Gð" ð 
ðA¸D¿IºIÀgÑ<NÔ<Nð Að AÝ4GðAð AÝWjðAð Aå*=ðAð Að Að	
õ Ø%Ø";Ø";Ø5Ø5Ø3Ø*KØ'EØ%AØ)Ø$Ð$Ñ&Ô&ðñ ô ð r&   ÚNonec                ó&   — t          | ¬¦  «         dS )u÷  Warm the `_load_deepgemm_kernel` cache from an opaque graph node, so Dynamo never traces the loader.

    Under `torch.compile`, Dynamo ignores `@functools.cache` and traces into `_load_deepgemm_kernel`,
    whose cold path (hub download + dynamic import via `lazy_load_kernel`) is untraceable and errors under
    `fullgraph`. `@allow_in_graph` turns the call into an opaque fx node instead â€” but an fx node's return
    must be proxyable, and the `DeepGEMM` bundle of Python callables isn't (`Unsupported: torch.* op
    returned non-Tensor`), so we can't just decorate the real loader. Hence two loaders: this one is
    opaque, returns `None`, and only warms the cache; the real `_load_deepgemm_kernel` right after is then
    a plain cache lookup.
    ©rX   N)ry   r|   s    r'   Ú_populate_deepgemm_kernelr}     s   € õ ¨Ð8Ñ8Ô8Ð8Ð8Ð8r&   c                óŽ   — t          | ¬¦  «         t          | ¬¦  «        }t          |t          ¦  «        rt	          |¦  «        ‚|S )Nr|   )r}   ry   Ú
isinstanceÚstrÚImportError)rX   Údeepgemm_or_errors     r'   Úload_deepgemm_kernelrƒ     sL   € Ý¨^Ð<Ñ<Ô<Ð<Ý-¸^ÐLÑLÔLÐÝÐ#¥SÑ)Ô)ð -ÝÐ+Ñ,Ô,Ð,ØÐr&   Údeviceútorch.devicec                óT   — t           j                             | ¦  «        d         dk    S )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@   )rn   r;   rp   ©r„   s    r'   Ú	_is_sm100rˆ      s$   € õ
 Œ:×+Ò+¨FÑ3Ô3°AÔ6¸"Ò<Ð<r&   Úscaleútorch.Tensorc                óz   — t          | j        ¦  «        sdS | j        t          j        k    rdS t          d¦  «        ‚)uº  On B200 (SM100) DeepGEMM only supports UE8M0 (power-of-two) scales; the float32 scales
    that work on H100 (SM90) have no SM100 path. UE8M0 scales load as ``float8_e8m0fnu`` (the
    loader normalizes even float32-container checkpoints like dsv4-flash-base), so a plain
    ``float32`` scale here means a genuine non-UE8M0 checkpoint â€” fail loud rather than let
    ``_coerce_sf_for_kernel`` silently round it and corrupt the output.
    NaÚ  DeepGEMM's Blackwell (SM100) experts kernel requires power-of-two (UE8M0) scale factors, but this checkpoint's expert scales are plain float32 (quantization_config.scale_fmt='float'). Rounding them to UE8M0 would scale the dequantized expert weights incorrectly and silently corrupt the output. Use a checkpoint quantized with scale_fmt='ue8m0', or an experts implementation that consumes float32 block scales directly, e.g. `model.set_experts_implementation('grouped_mm')`.)rˆ   r„   Údtypern   Úfloat32rH   )r‰   s    r'   Ú_assert_sm100_scales_are_ue8m0rŽ   (  sF   € õ �U”\Ñ"Ô"ð ØˆØ„{•e”mÒ#Ð#ØˆÝ
ð	<ñô ð r&   Úsfc                óª   — |                       t          j        ¦  «        }|dz                        d¦  «                              t          j        ¦  «        S )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  €ÿ)Úviewrn   Úint32Úbitwise_and_Úfloat)r�   Úint_views     r'   Ú_ceil_to_ue8m0r–   >  sA   € ð �wŠw•u”{Ñ#Ô#€HØ˜Ñ&×4Ò4Ð5EÑFÔF×KÒKÍEÌKÑXÔXÐXr&   Úexpected_mnú
int | Nonec                ó  — t          | j        ¦  «        }| j        t          j        k    r“|�H|                      d¦  «        |k     r/||                      d¦  «        z  }|                      |d¬¦  «        } |r2|                      ¦   «                              t          j	        ¦  «        } n;|  
                    ¦   «         } n&| j        t          j        k    r|rt          | ¦  «        } |                      ¦   «         dvr%t          d|                      ¦   «         › d�¦  «        ‚|s|                      ¦   «         S |                      d¦  «        }|                      d¦  «        }d|                      ¦   «         z  }| |z   |z  }|                      ¦   «         d	k    rd
|fn||z  d
|f}t!          |                      ¦   «         ¦  «        |k    r| S t          j        | j        || j        | j        ¬¦  «        }	|	                     | ¦  «         |	S )u¢  Lay out `sf` as DeepGEMM's dispatch expects, per arch.

    On SM100 the int-SF path only *checks* the SF (`tma_stride_check`) and never
    transforms it, so we hand it a TMA-aligned MN-major layout (`stride(-2) == 1`,
    `stride(-1) == align(mn, 16/esize)`). On SM90 DeepGEMM transforms SFA itself
    (`get_mn_major_tma_aligned_tensor`) and only *checks* SFB against
    `sm90_sfb_check`, which rejects TMA padding (`stride(-1)` must equal `size(-2)`,
    not `align(mn, â€¦)`); a padded weight SF trips `layout.hpp` whenever `mn` isn't a
    multiple of `16/esize` (e.g. N=576 â†’ mn=5). So on SM90 we return the raw
    row-major SF and let DeepGEMM lay it out.

    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.
    Néþÿÿÿ©Údim)r   r_   z"DeepGEMM SF must be 2D or 3D, got ÚDéÿÿÿÿé   r   r   ©rŒ   r„   )rˆ   r„   rŒ   rn   Úfloat8_e8m0fnuÚsizeÚrepeat_interleaveÚ
contiguousr‘   r’   r”   r�   r–   rœ   rH   Úelement_sizeÚtupleÚstrideÚempty_stridedÚshapeÚcopy_)
r�   r—   Úis_sm100Úgran_mnÚmnÚkfÚalign_toÚ
aligned_mnÚtarget_stridesÚouts
             r'   Ú_coerce_sf_for_kernelr³   K  sÎ  € õ8 ˜œÑ#Ô#€HØ	„x•5Ô'Ò'Ð'ØÐ" r§w¢w¨r¡{¤{°[Ò'@Ð'@Ø! R§W¢W¨R¡[¤[Ñ0ˆGØ×%Ò% g°2Ð%Ñ6Ô6ˆBØð 	Ø—’‘”×%Ò%¥e¤kÑ2Ô2ˆBˆBà—’‘”ˆBˆBØ	Œ•U”]Ò	"Ð	" xÐ	"Ý˜BÑÔˆà	‡v‚v�x„x�vÐÐÝÐI¸b¿fºf¹h¼hÐIÐIÐIÑJÔJÐJð ð Ø�}Š}‰ŒÐà	�Š�‰Œ€BØ	�Š�‰Œ€BØ�R—_’_Ñ&Ô&Ñ&€HØ�3˜(‘?Ð# hÑ.€JØ(*¯ª©¬°Aª¨�a˜�_�_¸BÀ¹OÈQÐPZÐ;[€NåˆR�YŠY‰[Œ[ÑÔ˜^Ò+Ð+Øˆ	Ý
Ô
˜bœh¨¸b¼hÈrÌyÐ
YÑ
YÔ
Y€CØ‡I‚Iˆb�M„M€MØ€Jr&   ÚweightÚweight_scale_invÚ
block_sizeútuple | Noner«   Údictc                óê   — | j         t          j        k    rddddœS |€t          d¦  «        ‚t	          |¦  «        }|dvrt          d|› d�¦  «        ‚|j         t          j        k    r|rdd	ddœS d
d	dœS )uC  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).
    Té    ©Ú	use_ue8m0Úgran_kÚuse_packed_ue8m0Nz]DeepGEMM requires block-wise quantized FP8 weights, but the experts have no `block_size` set.))é€   r¿   )r   r¿   u?   DeepGEMM requires `block_size` âˆˆ {(128, 128), (1, 128)}, got r>   r¿   F)r¼   r½   )rŒ   rn   Úint8rH   r¦   r¡   )r´   rµ   r¶   r«   s       r'   Ú_select_fp8_cast_kwargsrÁ   ˆ  s¨   € ð „|•u”zÒ!Ð!Ø!¨RÀTÐJÐJÐJàÐÝØkñ
ô 
ð 	
õ �zÑ"Ô"€JØÐ/Ð/Ð/ÝÐjÐ]gÐjÐjÐjÑkÔkÐkØÔ¥Ô!5Ò5Ð5¸(Ð5Ø!¨SÀdÐKÐKÐKØ¨#Ð.Ð.Ð.r&   Úexpert_ids_sortedÚnum_expertsr   Ú	alignmentÚuse_psum_layoutú&tuple[torch.Tensor, torch.Tensor, int]c                óÀ  — | j         }|                      d¦  «        }t          j        |                      ¦   «         |d|dz
  ¬¦  «                             ¦   «         }||z   dz
  |z  |z  }|t          ||¦  «        |dz
  z  z   }||z
  }	t          j        j         	                    |	 
                    d¦  «        d¦  «        }
t          j        ||¬¦  «        |
|          z   }|r(| 
                    d¦  «                             ¦   «         }nRt          j        |fd|t          j        ¬¦  «        }t          j        | |k     |                      ¦   «         d¦  «        ||<   |||fS )aŽ  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.
    r   r   )ÚbinsÚminÚmax)r   r   r‡   rž   ©r„   rŒ   )r„   r¢   rn   Úhistcr   ÚlongrÉ   ÚnnÚ
functionalÚpadÚcumsumÚarangeÚfullr’   Úwhere)rÂ   rÃ   rÄ   rÅ   r„   Ú
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Ý   §  sq  € ð Ô%€FØ"×'Ò'¨Ñ*Ô*€JåœÐ$5×$9Ò$9Ñ$;Ô$;À+ÐSTÐZeÐhiÑZiÐjÑjÔj×oÒoÑqÔqÐØ"3°iÑ"?À!Ñ"CÈ	Ñ!QÐU^Ñ ^Ðà"¥S¨°[Ñ%AÔ%AÀYÐQRÁ]Ñ%SÑSÐð 3Ð5FÑFÐÝœÔ,×0Ò0Ð1C×1JÒ1JÈ1Ñ1MÔ1MÈvÑVÔVÐÝ”| J°vÐ>Ñ>Ô>ÐASÐTeÔAfÑfÐàð uØ2×9Ò9¸!Ñ<Ô<×@Ò@ÑBÔBˆˆåœÐ%6Ð$8¸"ÀVÕSXÔS^Ð_Ñ_Ô_ˆÝ+0¬;Ð7HÈ;Ò7VÐXi×XmÒXmÑXoÔXoÐqsÑ+tÔ+tˆÐ'Ñ(à˜^Ð->Ð>Ð>r&   ÚxrÛ   rØ   c                ój   — t          j        |g| j        dd…         ¢R | j        | j        dœŽ}| ||<   |S )z;Pad a sorted tensor into the TMA-aligned contiguous layout.r   NrË   )rn   Úemptyr©   r„   rŒ   )rÞ   rÛ   rØ   Úpaddeds       r'   Ú_pad_for_deepgemmrâ   É  sD   € åŒ[Ð*ÐY¨Q¬W°Q°R°R¬[ÐYÐYÀÄÐQRÔQXÐYÐYÐY€FØ €FÐÑØ€Mr&   Úx_paddedc                ó   — | |         S ri   r%   )rã   rÛ   s     r'   Ú&_unpad_from_deepgemm_contiguous_layoutrå   Ð  s   € ØÐ$Ô%Ð%r&   Úhidden_statesÚtop_k_indexÚtop_k_weightsr   r¦   c                óx  — |                      d¦  «        }|                     d¦  «        }|                     d¦  «        }t          j        |¦  «        \  }	}
| |
|z           }||
         }t	          |	|||¦  «        \  }}}|	|k                         d¦  «        }|	                     |dz
  ¬¦  «         |||	||
|||fS )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Ê   )r¢   Úreshapern   Ú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ö   ×  sé   € ð × Ò  Ñ$Ô$€IØ×$Ò$ RÑ(Ô(€JØ"×*Ò*¨2Ñ.Ô.€Nõ œ JÑ/Ô/Ñ€L�$Ø*¨4°9Ñ+<Ô=ÐØ% dÔ+Ðõ
 ;\Ø�k ;°ñ;ô ;Ñ7Ð�nÐ&7ð " [Ò0×;Ò;¸BÑ?Ô?€MØ×Ò˜K¨!™OÐÑ,Ô,Ð,àØØØØØØØð	ð 	r&   Ú
out_paddedÚsorted_weightsrõ   rò   rÕ   rî   Ú
hidden_dimÚ	out_dtypeútorch.dtypec	                óÀ  — t          | |¦  «        }	|	|                     |	j        ¦  «                             d¦  «        z  }
|
                     |d¦  «         t          j        |¦  «        }t          j        |                     d¦  «        |	j	        ¬¦  «        ||<   |
|          
                    |||¦  «                             d¬¦  «                             |¦  «        S )uR   Unpad â†’ weighted multiply â†’ mask sentinels â†’ restore order â†’ top-k reduce.rž   g        r   r‡   r   r›   )rå   ÚtorŒ   rì   Úmasked_fill_rn   Ú
empty_likerÒ   r¢   r„   r‘   Úsum)r÷   rø   rõ   rò   rÛ   rÕ   rî   rù   rú   r²   ÚweightedÚinv_perms               r'   Ú_combine_routed_outputr  	  sÅ   € õ 1°Ð=MÑ
NÔ
N€CØ�^×&Ò& s¤yÑ1Ô1×;Ò;¸BÑ?Ô?Ñ?€Hð ×Ò˜-¨Ñ-Ô-Ð-ÝÔ Ñ%Ô%€HÝ”\ $§)¢)¨A¡,¤,°s´zÐBÑBÔB€HˆT�Nà�HÔ×"Ò" :¨y¸*ÑEÔE×IÒIÈaÐIÑPÔP×SÒSÐT]Ñ^Ô^Ð^r&   Úoutput_dtypezv5.16)r<   ÚinputÚbiasútorch.Tensor | Noneútorch.dtype | NoneÚactivation_scalec           
     óŒ  — |�t          d¦  «        ‚| j        t          j        t          j        fvrt          d| j        › �¦  «        ‚t          |j        t          j        k    ¬¦  «        }t          |||t          | j
        ¦  «        ¦  «        }|                      d| j        d         ¦  «        }	 |j        |	fi |¤Ž\  }
}t          j        |
j        d         |j        d         | j
        | j        ¬¦  «        }|                     d¦  «        rd	d	|d
         fnd}|                     |
t#          ||
                     d¦  «        ¬¦  «        f|t#          ||                     d¦  «        ¬¦  «        f||¬¦  «         |                     | j        dd…         |j        d         fz   ¦  «        }|�|                     |¦  «         |S )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.
    Nz@DeepGEMM linear does not support static activation quantization.z7DeepGEMM linear requires FP16 or BF16 activations, got r|   rž   r   rË   r¾   r   r½   ©r—   )Úrecipe)ÚNotImplementedErrorrŒ   rn   Úbfloat16Úfloat16rH   rƒ   rÀ   rÁ   rˆ   r„   r‘   r©   r   rà   r0   r   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Â  € ð Ð#Ý!Ð"dÑeÔeÐeØ„{�5œ>­5¬=Ð9Ð9Ð9ÝÐ`ÐSXÔS^Ð`Ð`ÑaÔaÐaå#°6´<Å5Ä:Ò3MÐNÑNÔN€HÝ)¨&Ð2BÀJÕPYÐZ_ÔZfÑPgÔPgÑhÔh€Kà�zŠz˜"˜eœk¨"œoÑ.Ô.€HØ8˜(Ô8¸ÐQÐQÀ[ÐQÐQÑ€IˆxÝŒ[˜œ¨Ô+¨V¬\¸!¬_ÀUÄ\ÐY^ÔYdÐeÑeÔe€Fð 2=·²ÐASÑ1TÔ1TÐ^��A�{ 8Ô,Ð-Ð-ÐZ^€IØ×ÒØ	Õ)¨(À	ÇÂÈqÑ@QÔ@QÐRÑRÔRÐSØ	Õ&Ð'7ÀVÇ[Â[ÐQRÁ^Ä^ÐTÑTÔTÐUØØð	 ñ ô ð ð �[Š[˜œ S b SÔ)¨V¬\¸!¬_Ð,>Ñ>Ñ?Ô?€FØÐØ�Š�DÑÔÐØ€Mr&   Úselfútorch.nn.Modulec                ó  — |j         t          j        k    rt          d|j         › �¦  «        ‚t	          ¦   «         }| j        r|j        n|j        }|j        }| 	                    d¦  «        }| 	                    d¦  «        }| 	                    d¦  «        }	t          |||| j        |j        t          |¦  «        ¦  «        \  }
}}}}}}}t          | j        r| j        n| j        ¦  «        }t          | j        ¦  «        }| j        r"t          | j        r| j        n| j        ¦  «        nd }| j        rt          | j        ¦  «        nd }| j        r|j        d         n|j        d         }t1          |
||¦  «        }t          j        ||||j         ¬¦  «        } |||||t          |¦  «        ¬¦  «         | j        r|                     d|||         ¦  «         | j        r|                      |¦  «        n|                      |¦  «        }t          j        ||	||j         ¬¦  «        } |||||t          |¦  «        ¬¦  «         | j        r|                     d|||         ¦  «         t;          ||||||||	|j         ¦	  «	        S )Nú;DeepGEMM experts path requires bfloat16 hidden states, got rž   r   r   rË   )rÅ   )rŒ   rn   r  rH   rƒ   Úis_transposedr   r   r„   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_matmulr„   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_forwardr2  M  s½  € ð Ô�eœnÒ,Ð,ÝÐlÐWdÔWjÐlÐlÑmÔmÐmå#Ñ%Ô%€Hà=AÔ=OÐt˜(Ô9Ð9ÐU]ÔUtÐàÔ!€FØ× Ò  Ñ$Ô$€IØ×#Ò# AÑ&Ô&€JØ×#Ò# BÑ'Ô'€Jõ 	Ø�{ M°4Ô3CÀXÔEYÕ[dÐekÑ[lÔ[lñ	ô 	ñ	ØØØØØØØØõ
 ¨d¬mÐM˜Ô*Ð*ÀÄÑNÔN€IÝ˜4œ>Ñ*Ô*€KØZ^ÔZgÐq�h°´ÐU�tÔ-Ð-ÀDÔDUÑVÔVÐVÐmq€GØ15´ÐH•˜Ô,Ñ-Ô-Ð-ÀD€Ið )-Ô(:ÐR�” Ô$Ð$À	ÄÐPQÔ@R€JÝ
˜MÐ+;Ð=NÑ
OÔ
O€CÝŒ{Ð,¨jÀÈ}ÔObÐcÑcÔc€HØÐ˜˜Y¨°.ÕR[Ð\bÑRcÔRcÐdÑdÔdÐdØ„}ð HØ×Ò˜AÐ/°¸Ô1FÑGÔGÐGà-1¬]ÐUˆt×Ò Ñ)Ô)Ð)ÀÇÂÈHÑ@UÔ@U€Hõ Œ+Ð'¨¸FÈ-ÔJ]Ð
^Ñ
^Ô
^€CØÐ˜ +¨s°NÕT]Ð^dÑTeÔTeÐfÑfÔfÐfØ„}ð EØ�Š�qÐ*¨I°lÔ,CÑDÔDÐDå!ØØØØØØØØØÔñ
ô 
ð 
r&   c                óÀ  — | j         rt          d¦  «        ‚t          | j        ¦  «         | j        dk    rt          d¦  «        ‚|j        t          j        k    rt          d|j        › �¦  «        ‚t          | j        j        t          j        k    ¬¦  «        }| j        r|j        n|j        }|j        }|                     d¦  «        }|                     d¦  «        }|                     d¦  «        }	t%          | j        r| j        n| j        ¦  «        }
t%          | j        r| j        n| j        ¦  «        }t%          | j        ¦  «        }t%          | j        ¦  «        }t1          |
|| j        t5          |¦  «        ¦  «        }t7          |||| j        |j        t5          |¦  «        ¦  «        \  }}}}}}}}|                     d¦  «        rd	d	|d
         fnd } |j        |fi |¤Ž\  }}tA          |||¦  «        }tA          |||¦  «        }t          j!        ||
j"        d	         |t          j        ¬¦  «        } ||tG          ||¬¦  «        f|
tG          ||
                     d¦  «        ¬¦  «        f|||t5          |¦  «        ¬¦  «         | j        r|  $                    |¦  «        n|  %                    |¦  «        } |j        |fi |¤Ž\  }}t          j!        ||	|t          j        ¬¦  «        } ||tG          ||¬¦  «        f|tG          ||                     d¦  «        ¬¦  «        f|||t5          |¦  «        ¬¦  «         tM          ||||||||	|j        ¦	  «	        S )NzðDeepGEMM experts selected on a model spanning multiple CUDA devices in one process; its kernels are bound to a single CUDA context and corrupt across devices. Use `experts_implementation='grouped_mm'`, or run one device per process (TP/EP).ÚstaticzJDeepGEMM experts dispatch does not support static activation quantization.r  r|   rž   r   r¾   r   r½   rË   r  rš   )r  rÅ   )'Ú_deepgemm_disabledÚRuntimeErrorrŽ   Údown_proj_scale_invÚactivation_schemer  rŒ   rn   r  rH   rƒ   r!  rÀ   r  r   r   r„   r¢   r   r  r  r   Úgate_up_proj_scale_invÚup_proj_scale_invrÁ   r¶   rˆ   rö   rÃ   r   r0   r   râ   rà   r©   r³   r'  r(  r  )r  ræ   rç   rè   r  Úgrouped_fp8_fp4_matmulr„   rî   rÕ   rù   r+  Úweight_scale_upr,  Úweight_scale_downr  r*  rø   Ú_expert_ids_grõ   rò   rÛ   rÜ   rØ   r  Úact_fp8Ú
act_scalesr1  Úproj_fp8Úproj_scalesr²   s                                 r'   Ú deepgemm_fp8_fp4_experts_forwardrC  Ž  sÈ  € ð Ôð 
õ ð\ñ
ô 
ð 	
õ # 4Ô#;Ñ<Ô<Ð<àÔ Ò)Ð)Ý!Ð"nÑoÔoÐoØÔ�eœnÒ,Ð,ÝÐlÐWdÔWjÐlÐlÑmÔmÐmå#°4´>Ô3GÍ5Ì:Ò3UÐVÑVÔV€Hà.2Ô.@ÐhˆÔ*Ð*ÀhÔFhð ð Ô!€FØ× Ò  Ñ$Ô$€IØ×#Ò# AÑ&Ô&€JØ×#Ò# BÑ'Ô'€Jå¨d¬mÐM˜Ô*Ð*ÀÄÑNÔN€IÝ¸d¼mÐg˜tÔ:Ð:ÐQUÔQgÑhÔh€OÝ˜4œ>Ñ*Ô*€KÝ  Ô!9Ñ:Ô:Ðå)¨)°_ÀdÄoÕW`ÐagÑWhÔWhÑiÔi€Kõ 	Ø�{ M°4Ô3CÀXÔEYÕ[dÐekÑ[lÔ[lñ	ô 	ñ	ØØØØØØØØð 2=·²ÐASÑ1TÔ1TÐ^��A�{ 8Ô,Ð-Ð-ÐZ^€Ið 9˜(Ô8¸ÐVÐVÈ+ÐVÐVÑ€GˆZÝ Ð)9Ð;LÑMÔM€GÝ" :Ð/?ÐARÑSÔS€JÝŒ{Ð,¨i¬o¸aÔ.@ÈÕW\ÔWeÐfÑfÔf€HØÐØ	Õ'¨
Ð@QÐRÑRÔRÐSØ	Õ)¨/ÀyÇ~Â~ÐVXÑGYÔGYÐZÑZÔZÐ[ØØØÝ! &Ñ)Ô)ðñ ô ð ð .2¬]ÐUˆt×Ò Ñ)Ô)Ð)ÀÇÂÈHÑ@UÔ@U€Hð ;˜HÔ:¸8ÐSÐSÀ{ÐSÐSÑ€HˆkÝ
Œ+Ð'¨¸FÍ%Ì.Ð
YÑ
YÔ
Y€CØÐØ	Õ(¨ÐBSÐTÑTÔTÐUØ	Õ+Ð,=È;×K[ÒK[Ð\^ÑK_ÔK_Ð`Ñ`Ô`ÐaØØØÝ! &Ñ)Ô)ðñ ô ð õ "ØØØØØØØØØÔñ
ô 
ð 
r&   Úmodulec                óP  — t          d¬¦  «        }t          | j        j        ¦  «        }t          | j        j        ¦  «        }t          | j        j        ¦  «                             t          j        ¦  «         	                    ¦   «         }t          | j
        j        ¦  «                             t          j        ¦  «         	                    ¦   «         }| j        }| j        }| j        }|dz  dk    s	|dz  dk    rt          d|› d|› d�¦  «        ‚|                     |                     ¦   «         d|z  |d	|¬
¦  «        }	|                     |                     ¦   «         ||d	|¬
¦  «        }
|                     ||	f||
f¦  «        \  \  }}	\  }}
t          j                             |d¬¦  «        | _        t          j                             |	d¬¦  «        | _        t          j                             |d¬¦  «        | _
        t          j                             |
d¬¦  «        | _        dS )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.
    Tr|   rº   r   zwDeepGEMM Mega MoE requires `hidden_dim` and `intermediate_hidden` divisible by 32 (FP8 SF granularity); got hidden_dim=z, intermediate_hidden=r>   r   )r   rº   )r  Ú
num_groupsF)Úrequires_gradN)rƒ   r   r9  Údatar7  r  r‘   rn   rÀ   r¤   r!  Úintermediate_dimrÃ   rù   rH   r   r”   r   rÎ   Ú	Parameter)rD  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_weightsrU  ê  s3  € õ" $°4Ð8Ñ8Ô8€HÝ˜fÔ;Ô@ÑAÔA€NÝ˜6Ô5Ô:Ñ;Ô;€Kå˜Ô,Ô1Ñ2Ô2×7Ò7½¼
ÑCÔC×NÒNÑPÔP€IÝ�fÔ&Ô+Ñ,Ô,×1Ò1µ%´*Ñ=Ô=×HÒHÑJÔJ€Fà Ô1ÐØÔ*ÐØÔ"€Jà�B�˜!ÒÐÐ2°RÑ7¸1Ò<Ð<ÝðmØ4>ðmð mØViðmð mð mñ
ô 
ð 	
ð
 ×;Ò;Ø×ÒÑÔØ	ÐÑØØØ$ð <ñ ô €Jð ×8Ò8Ø×ÒÑÔØØØØ$ð 9ñ ô €Gð .6×-TÒ-TØ	�JÐØ	�Ðñ.ô .Ñ*Ñ€Wˆj™?˜D 'õ  œ(×,Ò,¨WÀEÐ,ÑJÔJ€FÔÝ$)¤H×$6Ò$6°zÐQVÐ$6Ñ$WÔ$W€FÔ!Ý”x×)Ò)¨$¸eÐ)ÑDÔD€FÔÝ!&¤×!3Ò!3°GÈ5Ð!3Ñ!QÔ!Q€FÔÐÐr&   Úprocess_groupú%torch.distributed.ProcessGroup | Nonec                ób  — t          | j        ¦  «         | j        j        t          j        k    rt          d| j        j        › d�¦  «        ‚|€t          d¦  «        ‚t          d¬¦  «        }t          | dd¦  «        st          | ¦  «         d| _        |                     d	¦  «        }|                     d
¦  «        }|                     d	¦  «        }| j                             d
¦  «        }	| j                             d¦  «        dz  }
|	|                     ¦   «         z  }t          | dd¦  «        �| j        j        |k     r |                     ||||||
¬¦  «        | _        |                     |ddd¬¦  «        \  }}| j        j        d|…                              |¦  «         | j        j        d|…                              |¦  «         | j        j        d|…                              |¦  «         | j        j        d|…                              |¦  «         t	          j        ||ft          j        |j        ¬¦  «        }|                     || j        | j        f| j        | j        f| j        t          t          | dd¦  «        dd¦  «        ¬¦  «         |                     |j        ¦  «        S )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.Tr|   Ú_megamoe_transformedFrž   r   r   r   Úsymm_buffer)ÚhiddenÚnum_topkrÃ   Únum_max_tokens_per_rankrO  rº   r»   r    ÚconfigÚswiglu_limit)Úactivation_clamp)rŽ   r7  r  rŒ   rn   rÀ   r6  rH   rƒ   rq   rU  rY  r¢   rZ  r]  r   r   rÞ   rª   Úx_sfÚtopk_idxÚtopk_weightsrà   r  r„   r   r9  r!  rý   )r  ræ   rç   rè   rV  r  rî   rÕ   rù   rP  rO  Únum_global_expertsÚx_fp8ra  Úys                  r'   Ú(deepgemm_fp8_fp4_megamoe_experts_forwardrg  $  sØ  € õ2 # 4Ô#;Ñ<Ô<Ð<àÔÔ¥%¤*Ò,Ð,ÝðYØÔ!Ô'ðYð Yð Yñ
ô 
ð 	
ð
 ÐÝðiñ
ô 
ð 	
õ
 $°4Ð8Ñ8Ô8€Hõ �4Ð/°Ñ7Ô7ð )Ý˜dÑ#Ô#Ð#Ø$(ˆÔ!à× Ò  Ñ$Ô$€IØ×#Ò# AÑ&Ô&€JØ×#Ò# BÑ'Ô'€JØÔ)×.Ò.¨qÑ1Ô1ÐØÔ+×0Ò0°Ñ3Ô3°qÑ8ÐØ*¨]×-?Ò-?Ñ-AÔ-AÑAÐõ ˆt�] DÑ)Ô)Ð1°TÔ5EÔ5]Ð`jÒ5jÐ5jØ#×@Ò@ØØØØ*Ø$.Ø 3ð Añ 
ô 
ˆÔð ×0Ò0°È$ÐWYÐlpÐ0ÑqÔq�K€Eˆ4ØÔÔ�{˜
�{Ô#×)Ò)¨%Ñ0Ô0Ð0ØÔÔ˜+˜:˜+Ô&×,Ò,¨TÑ2Ô2Ð2ØÔÔ˜k˜z˜kÔ*×0Ò0°Ñ=Ô=Ð=ØÔÔ! + : +Ô.×4Ò4°]ÑCÔCÐCõ 	Œ�Z Ð,µE´NÈ=ÔK_Ð`Ñ`Ô`€AØ×ÒØ	Ø	Ô	˜DÔ7Ð8Ø	Œ˜Ô1Ð2ØÔÝ ¥¨¨x¸Ñ!>Ô!>ÀÐPTÑUÔUð ñ ô ð ð �4Š4�Ô#Ñ$Ô$Ð$r&   )r(   r)   )r(   r8   )F)rX   rY   r(   rZ   )rX   rY   r(   rz   )rX   rY   r(   r   )r„   r…   r(   rY   )r‰   rŠ   r(   rz   )r�   rŠ   r(   rŠ   ri   )r�   rŠ   r—   r˜   r(   rŠ   )
r´   rŠ   rµ   rŠ   r¶   r·   r«   rY   r(   r¸   )
rÂ   rŠ   rÃ   r   rÄ   r   rÅ   rY   r(   rÆ   )rÞ   rŠ   rÛ   rŠ   rØ   r   r(   rŠ   )rã   rŠ   rÛ   rŠ   r(   rŠ   )ræ   rŠ   rç   rŠ   rè   rŠ   rÃ   r   r   r   rÅ   rY   r(   r¦   )r÷   rŠ   rø   rŠ   rõ   rŠ   rò   rŠ   rÛ   rŠ   rÕ   r   rî   r   rù   r   rú   rû   r(   rŠ   )NNNN)r  rŠ   r´   rŠ   rµ   rŠ   r  r  r¶   r8   r  r  r	  r  r(   rŠ   )
r  r  ræ   rŠ   rç   rŠ   rè   rŠ   r(   rŠ   )rD  r  r(   rz   )r  r  ræ   rŠ   rç   rŠ   rè   rŠ   rV  rW  r(   rŠ   )6r#   Ú
__future__r   Ú	functoolsrD   r.   rJ   r1   Úcollections.abcr   Údataclassesr   rn   Úutilsr   Úutils.deprecationr   Úutils.import_utilsr	   r
   r   r   r   Úhub_kernelsr   Útensor_parallelr   Ú
get_loggerr    Úloggerr   Úcacher7   rW   ry   Ú_dynamoÚallow_in_graphr}   rƒ   rˆ   rŽ   r–   r³   rÁ   rÝ   râ   rå   rö   r  r  r2  rC  rU  rg  r%   r&   r'   ú<module>rv     s½  ðð
ð 
ð #Ð "Ð "Ð "Ð "Ð "à Ð Ð Ð Ø €€€Ø 	€	€	€	Ø 	€	€	€	Ø €€€Ø $Ð $Ð $Ð $Ð $Ð $Ø !Ð !Ð !Ð !Ð !Ð !à €€€à Ð Ð Ð Ð Ð Ø /Ð /Ð /Ð /Ð /Ð /ðð ð ð ð ð ð ð ð ð ð ð ð ð ð *Ð )Ð )Ð )Ð )Ð )Ø %Ð %Ð %Ð %Ð %Ð %ð 
ˆÔ	˜HÑ	%Ô	%€ð
 €�$ÐÑÔðð ð ð ð ñ ô ñ Ôðð, „ðð ð ñ „ðð$ „ð+ð +ð +ñ „ð+ð\ „ðpð pð pð pñ „ðpðf „Ôð9ð 9ð 9ð 9ñ Ôð9ðð ð ð ð ð „ð=ð =ð =ñ „ð=ðð ð ð ð,
Yð 
Yð 
Yð 
Yð:ð :ð :ð :ð :ðz/ð /ð /ð /ð>?ð ?ð ?ð ?ðDð ð ð ð&ð &ð &ð &ð/ð /ð /ð /ðd_ð _ð _ð _ð4 €�¨Ð1Ñ1Ô1ð
 !%Ø)-Ø'+Ø,0ð&ð &ð &ð &ñ 2Ô1ð&ðR>ð >ð >ð >ðBYð Yð Yð Yðx7Rð 7Rð 7Rð 7Rð~ <@ðS%ð S%ð S%ð S%ð S%ð S%ð S%r&   