
    ^j                    :   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j6                  e      Z ed       G d d             Zej>                  d(d       Z ej>                  d)d       Z!ej>                  d*d+d       Z"ejF                  jH                  d*d,d       Z%d*d-dZ&ej>                  d.d       Z'd/dZ(d0dZ)d1d2dZ*	 	 	 	 	 	 	 	 	 	 d3dZ+	 	 	 	 	 	 	 	 	 	 d4dZ,d5dZ-d6dZ.	 	 	 	 	 	 	 	 	 	 	 	 	 	 d7dZ/	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 d8dZ0 ed d!"      	 	 	 	 d9	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 d:d#       Z1	 	 	 	 	 	 	 	 	 	 d;d$Z2	 	 	 	 	 	 	 	 	 	 d;d%Z3d<d&Z4	 d1	 	 	 	 	 	 	 	 	 	 	 d=d'Z5y)>uS  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<   y)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__     m/var/www/ramen.bs-engineer-server.com/venv/lib/python3.12/site-packages/transformers/integrations/deepgemm.pyr   r   9   sI    H''''$$$$##'//$,,"** r&   r   c                 r   t         j                  j                  d      xs t         j                  j                  d      } | r| S t        j                  d      }|r<t         j
                  j                  t         j
                  j                  |            S t         j
                  j                  d      ryy)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_homer5   P   sy     

{+Jrzz~~k/JI<<Dwwrwwt455	ww}}&' r&   c                 *   t               } | yt        j                  j                  | d      }t        j                  j	                  |      r	 t        |      5 }t        j                  |      }ddd       j                  d|j                  di             j                  dd      }|j                  d      dd \  }}t        |      t        |      fS t        j                  j                  | d	      }t        j                  j	                  |      rp	 t        |      5 }t        j                  d
|j!                               }ddd       r4t        |j#                  d            t        |j#                  d            fS 	 t        j                  j                  | dd      }	t        j                  j	                  |	      rd	 t        |	      5 }t        j                  d|j!                               }ddd       r't        |j#                  d            }
|
dz  |
dz  dz  fS 	 yy# 1 sw Y   xY w# t        t        t        f$ r Y xw xY w# 1 sw Y   xY w# t        t        f$ r Y w xY w# 1 sw Y   xY w# t        t        f$ r Y yw xY w)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  
   )r5   r,   r1   joinisfileopenjsonloadr.   splitr   OSError
ValueErrorAttributeErrorresearchreadgroup)r4   version_jsonf
componentsr9   majorminorversion_txtmatchcuda_hcuda_versions              r'   _get_nvcc_versionrT   c   s>     I77<<	>:L	ww~~l#	l# *q!YYq\
* nn[*..2LMQQR[]_`G"==-bq1LE5u:s5z)) '',,y-8K	ww~~k"	k" Ja		">IJ5;;q>*CA,???  WW\\)Y9F	ww~~f	f K		"?JK"5;;q>2#t+lT.Ab-HHH  =* *
 ^4 		J J $ 		K K
 $ 		s   H7 H*/A&H7 I !%I=I J  %I460J  *H4/H7 7IIII I10I14I=9J   JJc                   t               sCt               sdt         dt         dt         dS t        j
                  j                         syt        j
                  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                  j                  t        j                  j                  |dd            sd| d|d    d|d    dS t               }|d| d|d    d|d    dS ||k  r/d| | d|d    d|d    d|d    d|d    d | d!|d    d|d    dS t        d"      }|y#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      }d$|	fd%|
fd&|fd'|fd(|fd)|fd+|fd,|fd-|fd.|fd/|ffD cg c]	  \  }}|| }}}|r*d0d1j                  |       d2t         dt         dt         d	S t#        |	|
|||||||| |       3      S c c}}w )4a  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=   )   rW   )rX      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   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	   torchr8   is_availableget_device_capabilityr5   r,   r1   r?   r>   rT   r   getattrr   r   )requires_sm100rN   rO   allowedarchmin_cudar4   nvcc_versionkernelr   r   r   r   r   r   r   r   r   get_mk_alignmentr   nameattrmissings                          r'   _load_deepgemm_kernelrp      s    $%#%g&'~6I5JJh&'q*
 zz&&(Nzz779u *%w*8&>bD'v-CE75'QRSS $rk7w"$	:8A;-qRS V5 5 ww~~bggll9eVDELYK X22:1+a}Lrt )**9+ 6%%-a[M8A;-yJ
 (" w.@!QxXY{m[n?#1\!_$5U9+ F$QK=(1+iA k*F~tV%6=N '0VX\ ] '0VX\ ]$V-PRVW$V-PRVW3FIfg(/8[]a(b%%,V5UW[%\"#*63QSW#X v'OQUVv'94@
 /35NO35NO02HI02HI*,AB02ST-/MN+-IJ57GH!12
D$ < 	G " ;DIIg<N;O P44G3HWjVk l**=)>aA	
 %";";553*K'E%A)$& 1s   I*c                    t        |        y)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.
    rf   N)rp   rr   s    r'   _populate_deepgemm_kernelrs     s     8r&   c                l    t        |        t        |       }t        |t              rt	        |      |S )Nrr   )rs   rp   
isinstancestrImportError)rf   deepgemm_or_errors     r'   load_deepgemm_kernelry     s2    ^<-^L#S)+,,r&   c                L    t         j                  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=   )rb   r8   rd   devices    r'   	_is_sm100r}      s#    
 ::++F3A6"<<r&   c                    t        | j                        sy| j                  t        j                  k7  ry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|   dtyperb   float32rE   )scales    r'   _assert_sm100_scales_are_ue8m0r   (  s8     U\\"{{emm#
	< r&   c                    | j                  t        j                        }|dz   j                  d      j                  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  )viewrb   int32bitwise_and_float)sfint_views     r'   _ceil_to_ue8m0r   >  s<     wwu{{#H&445EFKKEKKXXr&   c                   t        | j                        }| j                  t        j                  k(  r~|;| j                  d      |k  r'|| j                  d      z  }| j                  |d      } |r.| j                         j                  t        j                        } n;| j                         } n*| j                  t        j                  k(  r|rt        |       } | j                         dvrt        d| j                          d      |s| j                         S | j                  d      }| j                  d      }d| j                         z  }| |z   |z  }| j                         dk(  rd	|fn||z  d	|f}t!        | j#                               |k(  r| S t        j$                  | 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.
    dim)r   rY   z"DeepGEMM SF must be 2D or 3D, got D   r   r   r   r|   )r}   r|   r   rb   float8_e8m0fnusizerepeat_interleave
contiguousr   r   r   r   r   r   rE   element_sizetuplestrideempty_stridedshapecopy_)
r   expected_mnis_sm100gran_mnmnkfalign_to
aligned_mntarget_stridesouts
             r'   _coerce_sf_for_kernelr   K  s   8 #H	xx5'''"rwwr{['@!RWWR[0G%%g2%6B%%ekk2BB	U]]	"xB	vvxv=bffhZqIJJ }}	B	BR__&&H3(?#h.J(*Aa_BOQPZ;[NRYY[^+	


bhhbhhryy
YCIIbMJr&   c                    | 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_ue8m0z]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   rb   int8rE   r   r   )weightweight_scale_inv
block_sizer   s       r'   _select_fp8_cast_kwargsr     s     ||uzz!!RTJJk
 	
 z"J//\]g\hhijkk!5!55(!SdKK#..r&   c                   | j                   }| j                  d      }t        j                  | j	                         |d|dz
        j                         }||z   dz
  |z  |z  }|t        ||      |dz
  z  z   }||z
  }	t        j                  j                  j                  |	j                  d      d      }
t        j                  ||      |
|    z   }|r |j                  d      j	                         }nSt        j                  |fd|t        j                        }t        j                  | |k  | j	                         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   rb   histcr   longr   nn
functionalpadcumsumarangefullr   where)expert_ids_sortednum_experts	alignmentuse_psum_layoutr|   
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     sL    %%F"''*J$5$9$9$;+STZehiZijooq"3i"?!"C	!QU^ ^"S[%AYQR]%SS 35FF,,001C1J1J11MvV||Jv>ASTeAff299!<@@B%6$8"VSXS^S^_+0;;7H;7VXiXmXmXoqs+t'(^->>>r&   c                    t        j                  |g| j                  dd | j                  | j                  d}| ||<   |S )z;Pad a sorted tensor into the TMA-aligned contiguous layout.r   Nr   )rb   emptyr   r|   r   )xr   r   paddeds       r'   _pad_for_deepgemmr     sB    [[*YQWWQR[YQRQXQXYF FMr&   c                    | |   S Nr%   )x_paddedr   s     r'   &_unpad_from_deepgemm_contiguous_layoutr     s    $%%r&   c                <   |j                  d      }|j                  d      }|j                  d      }t        j                  |      \  }	}
| |
|z     }||
   }t	        |	|||      \  }}}|	|k\  j                  d      }|	j                  |dz
         |||	||
|||fS )zSort 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   reshaperb   sortr   	unsqueezeclamp_)hidden_statestop_k_indextop_k_weightsr   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$*49+<=%d+
 ;\k;;7n&7 "[0;;B?MK!O,	 	r&   c	                   t        | |      }	|	|j                  |	j                        j                  d      z  }
|
j	                  |d       t        j                  |      }t        j                  |j                  d      |	j                        ||<   |
|   j                  |||      j                  d      j                  |      S )uR   Unpad → weighted multiply → mask sentinels → restore order → top-k reduce.r   g        r   r{   r   r   )r   tor   r   masked_fill_rb   
empty_liker   r   r|   r   sum)
out_paddedsorted_weightsr   r   r   r   r   
hidden_dim	out_dtyper   weightedinv_perms               r'   _combine_routed_outputr   	  s     1=M
NC^&&syy1;;B??H --%H\\$))A,szzBHTNH"":y*EIIaIPSST]^^r&   output_dtypezv5.16)r9   c           
        |t        d      | j                  t        j                  t        j                  fvrt        d| j                         t        |j                  t        j                  k(        }t        |||t        | j                              }| j                  d| j                  d         }	 |j                  |	fi |\  }
}t        j                  |
j                  d   |j                  d   | j                  | j                        }|j                  d      rd	d	|d
   fnd}|j!                  |
t#        ||
j%                  d            f|t#        ||j%                  d            f||       |j                  | j                  dd |j                  d   fz         }||j'                  |       |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 rr   r   r   r   r   r   r   r   )recipe)NotImplementedErrorr   rb   bfloat16float16rE   ry   r   r   r}   r|   r   r   r   r   r.   r   r   r   add_)inputr   r   biasr   r   activation_scaledeepgemmcast_kwargsinput_2d	qinput_2dscale_2doutput	sf_recipes                 r'   deepgemm_fp8_fp4_linearr  #  s    #!"dee{{5>>5==99RSXS^S^R_`aa#6<<5::3MNH)&2BJPYZ_ZfZfPghKzz"ekk"o.H8(88Q[QIx[[+V\\!_U\\Y^YdYdeF 2=AS1TA{8,-Z^I	)(	q@QRS	&'7V[[QR^TU	   [[Sb)V\\!_,>>?FDMr&   c                b   |j                   t        j                  k7  rt        d|j                          t	               }| j
                  r|j                  n|j                  }|j                  }|j                  d      }|j                  d      }|j                  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        j2                  ||||j                         } |||||t        |             | j&                  r|j5                  d|||          | j                  r| j7                  |      n| j9                  |      }t        j2                  ||	||j                         } |||||t        |             | j&                  r|j5                  d|||          t;        ||||||||	|j                   	      S )N;DeepGEMM experts path requires bfloat16 hidden states, got r   r   r   r   )r   )r   rb   r   rE   ry   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   )selfr   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_forwardr  M  sj    enn,VWdWjWjVklmm#%H=A=O=O(99U]UtUt!!F  $I##A&J##B'J 	{M43C3CXEYEY[dek[l		
 dmm**NI4>>*KZ^ZgZght--DDUDUVmqG15,,-DI )-(:(:$	PQ@RJ
M+;=N
OC{{,j}ObObcHY.R[\bRcd}}A/1FG-1]]t)H@UH ++'F-J]J]
^C+sNT]^dTef}}q*Il,CD!
 
r&   c                   | j                   rt        d      t        | j                         | j                  dk(  rt        d      |j                  t        j                  k7  rt        d|j                         t        | j                  j                  t        j                  k(        }| j                  r|j                  n|j                  }|j                   }|j#                  d      }|j#                  d      }|j#                  d      }	t%        | j&                  r| j(                  n| j*                        }
t%        | j&                  r| j,                  n| j.                        }t%        | j                        }t%        | j                        }t1        |
|| j2                  t5        |            }t7        |||| j8                  |j:                  t5        |            \  }}}}}}}}|j=                  d      rd	d	|d
   fnd } |j>                  |fi |\  }}tA        |||      }tA        |||      }t        jB                  ||
jD                  d	   |t        j                        } ||tG        ||      f|
tG        ||
j#                  d            f|||t5        |             | j&                  r| jI                  |      n| jK                  |      } |j>                  |fi |\  }}t        jB                  ||	|t        j                        } ||tG        ||      f|tG        ||j#                  d            f|||t5        |             tM        ||||||||	|j                  	      S )NzDeepGEMM 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  rr   r   r   r   r   r   r   r   r   )r   r   )'_deepgemm_disabledRuntimeErrorr   down_proj_scale_invactivation_schemer   r   rb   r   rE   ry   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   r.   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_scalesr  proj_fp8proj_scalesr   s                                 r'    deepgemm_fp8_fp4_experts_forwardr0    sG     \
 	
 #4#;#;<)!"nooenn,VWdWjWjVklmm#4>>3G3G5::3UVH.2.@.@**hFhFh  !!F  $I##A&J##B'Jdmm**NIdmmt::QUQgQghO4>>*K !9!9:))_dooW`agWhiK 	{M43C3CXEYEY[dek[l		 2=AS1TA{8,-Z^I 9(88V+VGZ)9;LMG":/?ARSJ{{,iooa.@W\WeWefH	'
@QRS	)/y~~VXGYZ[!&) .2]]t)H@UH ;H::8S{SHk
++'F%..
YC	(BSTU	+,=;K[K[\^K_`a!&) "
 
r&   c                P   t        d      }t        | j                  j                        }t        | j                  j                        }t        | j
                  j                        j                  t        j                        j                         }t        | j                  j                        j                  t        j                        j                         }| j                  }| j                  }| j                  }|dz  dk7  s|dz  dk7  rt        d| d| d      |j                  |j!                         d|z  |d	|
      }	|j                  |j!                         ||d	|
      }
|j#                  ||	f||
f      \  \  }}	\  }}
t        j$                  j'                  |d      | _        t        j$                  j'                  |	d      | _        t        j$                  j'                  |d      | _
        t        j$                  j'                  |
d      | _        y)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.
    Trr   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)ry   r   r&  datar$  r  r   rb   r   r   r  intermediate_dimr   r   rE   r   r   r   r   	Parameter)moduler   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_weightsrB    s   " $48Hf;;@@AN655::;K,,11277

CNNPIf&&++,11%**=HHJF 11**""JB!2R71<44><?UViUjjkm
 	

 ;;	$ < J 88$ 9 G .6-T-T	J	.*Wj?D'  ((,,WE,JF$)HH$6$6zQV$6$WF!xx))$e)DF!&!3!3G5!3!QFr&   c                   t        | j                         | j                  j                  t        j
                  k7  r#t        d| j                  j                   d      |t        d      t        d      }t        | dd      st        |        d| _        |j                  d	      }|j                  d
      }|j                  d	      }| j                  j                  d
      }	| j                  j                  d      dz  }
|	|j                         z  }t        | dd      | j                  j                  |k  r|j                  ||||||
      | _        |j!                  |ddd      \  }}| j                  j"                  d| j%                  |       | j                  j&                  d| j%                  |       | j                  j(                  d| j%                  |       | j                  j*                  d| j%                  |       t	        j,                  ||ft        j.                  |j0                        }|j3                  || j                  | j4                  f| j6                  | j                  f| j                  t        t        | dd      dd             |j9                  |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.NzDeepGEMM Mega MoE requires a `process_group` for the EP group. The TP wrapping (MoeTensorParalellMegaMoeExperts) supplies it automatically; pass it explicitly otherwise.Trr   _megamoe_transformedFr   r   r   r   symm_buffer)hiddennum_topkr   num_max_tokens_per_rankr<  r   r   r   configswiglu_limit)activation_clamp)r   r$  r  r   rb   r   r#  rE   ry   re   rB  rD  r   rE  rH  r   r   r   r   x_sftopk_idxtopk_weightsr   r   r|   r   r&  r  r   )r  r   r   r   process_groupr   r   r   r   r=  r<  num_global_expertsx_fp8rL  ys                  r'   (deepgemm_fp8_fp4_megamoe_experts_forwardrS  $  s   2 #4#;#;<%**,!!''((WY
 	

 i
 	

 $48H 4/7d#$(!  $I##A&J##B'J))..q1++003q8*]-?-?-AA t]D)1T5E5E5]5]`j5j#@@*$. 3 A 
 00$WYlp0qKE4{
#))%0+:&,,T2kz*00=!!+:.44]C 	Z,ENN=K_K_`A				D778	112 x!>PTU   44##$$r&   )returnz
str | None)rT  tuple[int, int] | None)F)rf   boolrT  zDeepGEMM | str)rf   rV  rT  None)rf   rV  rT  r   )r|   ztorch.devicerT  rV  )r   torch.TensorrT  rW  )r   rX  rT  rX  r   )r   rX  r   z
int | NonerT  rX  )
r   rX  r   rX  r   ztuple | Noner   rV  rT  dict)
r   rX  r   r   r   r   r   rV  rT  z&tuple[torch.Tensor, torch.Tensor, int])r   rX  r   rX  r   r   rT  rX  )r   rX  r   rX  rT  rX  )r   rX  r   rX  r   rX  r   r   r   r   r   rV  rT  r   )r   rX  r   rX  r   rX  r   rX  r   rX  r   r   r   r   r   r   r   ztorch.dtyperT  rX  )NNNN)r   rX  r   rX  r   rX  r   torch.Tensor | Noner   rU  r   ztorch.dtype | Noner   rZ  rT  rX  )
r  torch.nn.Moduler   rX  r   rX  r   rX  rT  rX  )r7  r[  rT  rW  )r  r[  r   rX  r   rX  r   rX  rO  z%torch.distributed.ProcessGroup | NonerT  rX  )6r#   
__future__r   	functoolsrA   r,   rG   r/   collections.abcr   dataclassesr   rb   utilsr   utils.deprecationr   utils.import_utilsr	   r
   r   r   r   hub_kernelsr   tensor_parallelr   
get_loggerr    loggerr   cacher5   rT   rp   _dynamoallow_in_graphrs   ry   r}   r   r   r   r   r   r   r   r   r   r  r  r0  rB  rS  r%   r&   r'   <module>rj     s[  
 #   	 	  $ !   /  * % 
		H	%
 $  ,  $ + +\ p pf 9 9 = =,
Y:z//,8/FR/^b/	/>?#?25?BE?X\?+?D&///  / 	/
 / / /d__ _  _ 	_
 #_ _ _ _ _ _4 1
 !%)-'+,0&&& #& 	&
 '& %& *& & 2&R>
>> >  	>
 >BY
YY Y  	Y
 Yx7R~ <@S%
S%S% S%  	S%
 9S% S%r&   