
    ^jj                       d dl mZ d dlmZ d dlmZ ddlmZ ddlm	Z	 ddl
mZmZmZmZ dd	lmZ dd
lmZ  e       r:d dlZej(                  j+                  e      Zej(                  j+                  e      Z ej,                  e      Z	 	 d"	 	 	 	 	 	 	 	 	 d#dZ	 	 	 	 	 	 	 	 	 	 d$dZd%dZd%dZd Zd Z e       rYej>                  jA                  dedd       ej>                  jC                  de       ej>                  jE                  dee       d&dZ#	 	 	 	 	 	 	 	 d%dZ$	 	 d"	 	 	 	 	 	 	 	 	 	 	 d'dZ%	 	 	 	 	 	 	 	 	 	 d$dZ& G d de	      Z' e'       Z(d(dZ)	 d)e(ddddd 	 	 	 	 	 	 	 	 	 	 	 	 	 d*d!Z*y)+    )annotations)Callable)wraps   )logging)GeneralInterface)is_torch_availableis_torch_greater_or_equalis_torch_less_or_equalis_torchdynamo_compiling   )deepgemm_bf16_experts_forward)sonicmoe_experts_forwardNFc                   |r5t        j                  | j                  d      |      j                  d      }n4t        j                  || j                  d            j                  d      }||j	                  |       |S )a  Batched linear layer supporting optional bias and transposed weights.

    Args:
        input (`torch.Tensor`):
            Input tensor of shape (batch_size, input_dim).
        weight (`torch.Tensor`):
            Weight tensor of shape (batch_size, output_dim, input_dim) if transposed is `False`,
            else of shape (batch_size, input_dim, output_dim).
        bias (`torch.Tensor`, *optional*):
            Bias tensor of shape (batch_size, output_dim). Default is `None`.
        is_transposed (`bool`, *optional*, defaults to `False`):
            Whether the weight tensor is transposed.
    Returns:
        `torch.Tensor`: Output tensor of shape (batch_size, output_dim).
    r   )torchbmm	unsqueezesqueezeadd_)inputweightbiasis_transposedouts        h/var/www/ramen.bs-engineer-server.com/venv/lib/python3.12/site-packages/transformers/integrations/moe.py_batched_linearr   T   sg    * ii*F3;;A> ii 34<<R@J    c                   |j                  d      }|j                  d      }|j                  d      }|j                  |d      }|j                  d      }|j                  d      }	|	j                  d| j                  dz
        }	| j
                  r-| j                  |	   }
| j                  r| j                  |	   nd }n,| j                  |	   }
| j                  r| j                  |	   nd }t        ||
|| j                        }| j
                  r| j                  |      }n| j                  |      }| j                  |	   }
| j                  r| j                   |	   nd }t        ||
|| j                        }||j#                  d      z  }|j%                  |||      j'                  d      }|j)                  |j*                        S )Nr   r   dimr   r   r   )sizerepeat_interleavereshapeclampnum_expertshas_gategate_up_projhas_biasgate_up_proj_biasup_projup_proj_biasr   r   _apply_gateact_fn	down_projdown_proj_biasr   viewsumtodtype)selfhidden_statestop_k_indextop_k_weights	num_top_k
num_tokens
hidden_dimselected_hidden_statessample_weights
expert_idsselected_weightsselected_biasesproj_outweighted_outfinal_hidden_statess                  r   batched_mm_experts_forwardrE   v   s      $I##A&J##B'J +<<YA<N"**2.N$$R(J !!!T%5%5%9:J }},,Z8@D$00<SW<<
3;?==$++J7d  0VZVhVhH
 }}##H- ;;x( ~~j19=d))*5DO "HZHZH
 n66r::L '++J	:NRRWXRY!!-"5"566r   c                4   t        j                  | j                  d      |j                  d      | j                  | j                        }d}t        |j                               D ].  \  }}||k(  rt        j                  | || ||   |||        |}0 |S )a(  
    Fallback grouped matrix multiplication used when `torch.nn.functional.grouped_mm` and `torch._grouped_mm`
    are unavailable or incompatible with `torch.compile` (e.g. non-bfloat16 weights).

    Args:
        input (`torch.Tensor`): Input of shape (S, input_dim), sorted by expert id.
        weight (`torch.Tensor`): Expert weights of shape (num_experts, input_dim, output_dim).
        offs (`torch.Tensor`): Cumulative token counts per expert of shape (num_experts,).
    Returns:
        `torch.Tensor`: Output of shape (S, output_dim).
    r   r   devicer5   r   )r   zerosr#   rH   r5   	enumeratetolistmm)r   r   offsoutputstartiends          r   _grouped_mm_fallbackrS      s     [[AAu||SXS^S^_FE DKKM* 3C<uS!6!9&s2CD	 Mr   c                p   | j                         dk(  sJ dt        | j                                |j                         dk(  sJ dt        |j                                |j                         dk(  sJ dt        |j                                |j                  d      |j                  d      k(  s+J d|j                  d       d	|j                  d              | j                  d      |j                  d      k(  s+J d
| j                  d       d|j                  d              |j                  t
        j                  t
        j                  fv sJ d|j                          t        j                  | j                  d      |j                  d      | j                  | j                        S )zRShape/dtype inference stub for `_grouped_mm_fallback` required by `torch.compile`.r   z+input must be 2D (S, input_dim), got shape    zBweight must be 3D (num_experts, input_dim, output_dim), got shape r   z*offs must be 1D (num_experts,), got shape r   zoffs length z must match number of experts zinput_dim mismatch: input has z, weight has z$offs must be an integer tensor, got rG   )
r!   tupleshaper#   r5   r   int32int64emptyrH   r   r   rN   s      r   _grouped_mm_fallback_faker\      s   99;!_J5QVQ\Q\K]J^__::<1 
LUSYS_S_M`Lab 88:?\HtzzIZH[\\?99Q<6;;q>)v\$))A,Geflfqfqrsfteu+vv)::a=FKKN* 
(A}V[[QR^DTU* ::%++u{{33h7[\`\f\f[g5hh3;;uzz!}fkk!nU\\QVQ\Q\]]r   c                H    | j                  |d   |d          |d   | _        y)zjSaves input and weight for backward; offs is stored directly as it is a non-differentiable integer tensor.r   r   r   N)save_for_backwardrN   )ctxinputsrO   s      r   "_grouped_mm_fallback_setup_contextra      s%    &)VAY/ayCHr   c                   | j                   \  }}t        j                  |      }t        j                  |      }d}t        | j                  j                               D ]c  \  }}||k(  rt        j                  ||| ||   j                  |||        t        j                  ||| j                  ||| ||          |}e ||dfS )zuBackward pass for `_grouped_mm_fallback`. Computes grad_input and grad_weight per expert group; offs has no gradient.r   rI   N)saved_tensorsr   
zeros_likerK   rN   rL   rM   T)	r_   grad_outputr   r   
grad_inputgrad_weightrP   rQ   rR   s	            r   _grouped_mm_fallback_backwardri      s    %%ME6!!%(J""6*KE CHHOO-. 3C<U3'*U3:OPuS!##[s%;QP {D((r   z!transformers::grouped_mm_fallback z4(Tensor input, Tensor weight, Tensor offs) -> Tensor)mutates_argsschema)setup_contextc                l   t               r|j                  t        j                  k7  sx|j                  j
                  dk(  r9t        dd      r,|j                         dz  dk7  s<| j                         dz  dk7  s&|j                  j
                  dk(  rt        dd      ry|j                  j
                  d	k(  rt        t        j                  j                  d
      r,t        j                  j                  |j                        dk\  S t        t        d      ret        dd      r,t        j                  j                  |j                        dk\  S t        j                  j                  |j                        dk\  S yt        t        j                  j                  d
      xs t        t        d      S )a  
    Check if torch.nn.functional.grouped_mm or torch._grouped_mm can be used based on availability and compatibility with torch.compile.

    Args:
        input (`torch.Tensor`):
            Input tensor of shape (S, input_dim).
        weight (`torch.Tensor`):
            Weight tensor of shape (num_experts, input_dim, output_dim).
        offs (`torch.Tensor`):
            Offsets tensor indicating the boundaries of each group in the input tensor.
    Returns:
        `bool`: True if grouped_mm can be used, False otherwise.
    cpuz2.10.0T)
accept_dev   r   z2.8.0Fcuda
grouped_mm)   r   _grouped_mmz2.9)	   r   )r   r5   r   bfloat16rH   typer   data_ptrhasattrnn
functionalrr   get_device_capabilityr
   r[   s      r   _can_use_grouped_mmr~   
  s9     
"	#(F==&"8=__#q(ENN,<r,AQ,F==&"7t< 
 }}V#588&&5::33FMMBfLL5-((4@zz77F&PPzz77F&PP588&&5V9VVr   c                   t        | ||      rt        t        j                  j                  d      rEt        j                  j                  j                  | j                  |j                        ||      S t        t        d      r1t        j                  | j                  |j                        ||      S t        j                  j                  j                  | ||      S )a  Grouped matrix multiplication dispatcher that uses torch.nn.functional.grouped_mm if available, else falls back to torch._grouped_mm.

    Args:
        input (`torch.Tensor`):
            Input tensor of shape (S, input_dim).
        weight (`torch.Tensor`):
            Weight tensor of shape (num_experts, input_dim, output_dim).
        offs (`torch.Tensor`):
            Offsets tensor indicating the boundaries of each group in the input tensor.
    Returns:
        `torch.Tensor`: Output tensor of shape (S, output_dim).
    rs   rN   ru   )r~   rz   r   r{   r|   rs   r4   r5   ru   opstransformersgrouped_mm_fallbackr[   s      r   ru   ru   :  s    $ 5&$/
 588&&588&&11%((6<<2H&W[1\\UM*$$UXXfll%;V$OO99!!55eV$5OOr   c                    |rt        | ||      }nt        | |j                  dd      |      }||j                  |       |S )a  Grouped linear layer supporting optional bias and transposed weights.

    Args:
        input (`torch.Tensor`):
            Input tensor of shape (S, input_dim).
        weight (`torch.Tensor`):
            Weight tensor of shape (num_experts, input_dim, output_dim) if `is_transposed`,
            else of shape (num_experts, output_dim, input_dim).
        offs (`torch.Tensor`):
            Offsets tensor indicating the boundaries of each group in the input tensor.
        bias (`torch.Tensor`, *optional*):
            Bias tensor of shape (num_experts, output_dim). Default is `None`.
        is_transposed (`bool`, *optional*, defaults to `False`):
            Whether the weight tensor is transposed.
    Returns:
        `torch.Tensor`: Output tensor of shape (S, output_dim).
    r   r   )ru   	transposer   )r   r   rN   r   r   r   s         r   _grouped_linearr   Y  sH    0 %d3 %!1!1"b!9EJr   c                   |j                   }|j                  d      }|j                  d      }|j                  d      }|j                  d      }|j                  d      }	t        j                  |	      \  }
}|||z     }||   }|j
                  dv r|
j                         n|
j                         }t        j                  || j                  d| j                  dz
        }t        j                  |dt        j                        }|
| j                  k\  j                  d      }|
j                  | j                  dz
         | j                  r*| j                  }| j                   r| j"                  |
   nd }n)| j$                  }| j                   r| j&                  |
   nd }|j)                  |d       t+        ||||| j,                  	      }| j                  r| j/                  |      }n| j1                  |      }| j2                  }| j                   r| j4                  |
   nd }t+        ||||| j,                  	      }||j                  d      z  }|j)                  |d       t        j6                  |      }t        j8                  |j                  d      |
      ||<   ||   }|j;                  |||      j=                  d      }|j?                  |j@                        S )Nr   r   )ro   mpsr   )binsminmax)r!   r5   )r   g        r"   )rH   r    )!rH   r#   r%   r   sortrx   floatinthistcr'   cumsumrX   r   clamp_r(   r)   r*   r+   r,   r-   masked_fill_r   r   r.   r/   r0   r1   
empty_likearanger2   r3   r4   r5   )r6   r7   r8   r9   rH   r:   r;   r<   r>   r?   expert_ids_gpermselected_hidden_states_gsample_weights_ghistc_inputtokens_per_expertoffsetssentinel_maskr@   rA   rB   rC   inv_permrD   s                           r   grouped_mm_experts_forwardr     s    !!F  $I##A&J##B'J #**2.N$$R(J J/L$,TY->?%d+ +1++*G,$$&\M]M]M_KKd6F6FASWScScfgSghll,!5;;GG" "T%5%55@@DMD,,q01 }},,BF--$00>UY<<=A]]$++L9PT ))-=  "2G/aeasasH
 }}##H- ;;x( ~~;?==d)),7dO "G/QUQcQcH
 .88<<L mS1 %H\\$))A,v>HTN)L '++J	:NRRWXRY!!-"5"566r   c                  2     e Zd ZdZeeeedZd fdZ	 xZ
S )ExpertsInterfacez;Interface for registering custom experts forward functions.)deepgemm
batched_mmrs   sonicmoec                    |t         j                  d       n|dk7  r|| vrt        d| d      t        |   ||      S )zfReturn the requested `experts_implementation`. Also strictly check its validity, and raise if invalid.a
  You tried to access the `ExpertsInterface` with a `config._experts_implementation` set to `None`. This is expected if you use an Expert Module as a standalone Module. If this is not the case, something went wrong with the dispatch of `config._experts_implementation`eager`zL` is not a valid experts implementation registered in the `ExpertsInterface`)loggerwarning_onceKeyErrorsuperget)r6   experts_implementationdefault	__class__s      r   get_interfacezExpertsInterface.get_interface  s`    !)N
 $w.3IQU3U*++wx  w{17;;r   )r   strr   r   returnr   )__name__
__module____qualname____doc__r   rE   r   r   _global_mappingr   __classcell__)r   s   @r   r   r     s%    E 200,	O< <r   r   c                V    |j                  dd      \  }}| j                  |      |z  S )a  
    Default gating mechanism: splits the gate_up_out into gate and up parts,
    applies the activation function to the gate part, and multiplies it with the up part.
    Args:
        gate_up_out (`torch.Tensor`):
            The output tensor from the gate and up projection of shape (S, 2 * intermediate_dim).
    Returns:
        `torch.Tensor`: The gated output tensor of shape (S, intermediate_dim).
    r   r   r    )chunkr/   )r6   gate_up_outgateups       r   _default_apply_gater     s1        +HD";;tr!!r   T)experts_interfaceis_concatenatedr   r*   r(   c               8    dfd}|  ||       S |S )a  Decorator to modify experts class to support different experts implementations.

    Args:
        experts_class (`type[torch.nn.Module]`, *optional*):
            The experts class to modify. If not provided, returns a decorator that can be applied to the class.
        experts_interface (`ExpertsInterface`, *optional*, defaults to `ALL_EXPERTS_FUNCTIONS`):
            The experts interface to use for dispatching the forward method.
        is_concatenated (`bool`, *optional*, defaults to `True`):
            Whether the expert weights are stored in concatenated layout [gate;up]
            or interleaved layout [gate0, up0, gate1, up1, ...].
        is_transposed (`bool`, *optional*, defaults to `False`):
            Whether the expert weights are stored in transposed format.
        has_bias (`bool`, *optional*, defaults to `False`):
            Whether the expert layers include bias terms or not.
        has_gate (`bool`, *optional*, defaults to `True`):
            Whether the experts use a gating mechanism or not.
            Whether it has gate_up_proj weights or just up_proj weights.

    Returns:
        `type[torch.nn.Module]`: The modified experts class.
    c                    | j                   | j                  t              	fd       }t              fd       }t        | d      st        | _        || _         || _        | S )Nc                f     | |g|i | || _         | _        | _        | _        | _        y N)configr(   r*   r   r   )	r6   r   argskwargsr*   r(   r   r   original_inits	       r   __init__z=use_experts_implementation.<locals>.wrapper.<locals>.__init__4  s<    $888 DK$DM$DM!.D#2D r   c                h    j                  | j                  j                        } || g|i |S r   )r   r   _experts_implementation)r6   r   r   experts_forwardr   original_forwards       r   forwardz<use_experts_implementation.<locals>.wrapper.<locals>.forward=  s5    /==dkk>a>acstO"49$9&99r   r.   )r   r   r   rz   r   r.   )
experts_classr   r   r   r   r   r*   r(   r   r   s
      @@r   wrapperz+use_experts_implementation.<locals>.wrapper0  su    %..(00	}		3 
	3 
	 	: 
!	: }m4(;M%!) 'r   )r   type[torch.nn.Module]r   r   rj   )r   r   r   r   r*   r(   r   s    ````` r   use_experts_implementationr     s%    > 2  }%%Nr   )NF)
r   torch.Tensorr   r   r   torch.Tensor | Noner   boolr   r   )
r6   ztorch.nn.Moduler7   r   r8   r   r9   r   r   r   )r   r   r   r   rN   r   r   r   )r   r   r   r   rN   r   r   r   )r   r   r   r   rN   r   r   r   r   r   r   r   )r   r   r   r   r   )r   ztype[torch.nn.Module] | Noner   r   r   r   r   r   r*   r   r(   r   r   r   )+
__future__r   collections.abcr   	functoolsr   utilsr   utils.genericr   utils.import_utilsr	   r
   r   r   r   r   r   r   r   _dynamoassume_constant_result
get_loggerr   r   r   rE   rS   r\   ra   ri   library	custom_opregister_fakeregister_autogradr~   ru   r   r   r   ALL_EXPERTS_FUNCTIONSr   r   rj   r   r   <module>r      s   # $   ,  4 . 
 !& D DE^ _"]]AABXY 
		H	%\ !%	  	
 D=7
=7=7 =7  	=7
 =7F4^)& 	MM+E	   
MM CE^_	MM##+%8 $ -W`PPP P 	PF !%### # 	#
 # #Le7
e7e7 e7  	e7
 e7P<' <2 )* " 37; +@ ;/; (; 	;
 ; ; ; ;r   