
    ^jqW                        d Z ddlmZ ddlmZ ddlZddlmc mZ	 ddlm
Z
mZmZ ddlmZ ddlmZ dd	lmZ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mZ ddl m!Z!m"Z" ddl#m$Z$m%Z% ddl&m'Z' ddl(m)Z)m*Z*m+Z+m,Z, ddl-m.Z.m/Z/ ddl0m1Z1 ddl2m3Z3  e,jh                  e5      Z6d Z7 G d dejp                        Z9 G d dejt                        Z;d Z<d=dZ= G d dejt                        Z>dej                  d e?d!ej                  fd"Z@	 d>d#ejt                  d$ej                  d%ej                  d&ej                  d'ej                  dz  d(eAd)eAd*e'e)   fd+ZB G d, d-ejt                        ZC G d. d/e      ZDe* G d0 d1e%             ZEe* G d2 d3eE             ZF G d4 d5eEe      ZG G d6 d7eeE      ZH G d8 d9eeE      ZI G d: d;eeE      ZJg d<ZKy)?zPyTorch Nemotron model.    )Callable)OptionalN)SizeTensornn   )initialization)ACT2FN)CacheDynamicCache)GenerationMixin)create_causal_mask)GenericForQuestionAnswering GenericForSequenceClassificationGenericForTokenClassificationGradientCheckpointingLayer)BaseModelOutputWithPastCausalLMOutputWithPast)ROPE_INIT_FUNCTIONSdynamic_rope_update)ALL_ATTENTION_FUNCTIONSPreTrainedModel)Unpack)TransformersKwargsauto_docstringcan_return_tuplelogging)maybe_autocastmerge_with_config_defaults)capture_outputs   )NemotronConfigc                     t        j                         s|S t        j                  |       }t         j                  j                  j                  || |      S N)torchis_autocast_enabledget_autocast_dtypeampautocast_mode_cast)device_typeargstarget_dtypes      y/var/www/ramen.bs-engineer-server.com/venv/lib/python3.12/site-packages/transformers/models/nemotron/modeling_nemotron.py_cast_if_autocast_enabledr/   6   sB    $$&//<yy&&,,T;MM    c            	       ^     e Zd Z	 	 	 	 	 d	deee   z  ez  dededef fdZde	de	fdZ
 xZS )
NemotronLayerNorm1Pnormalized_shapeepselementwise_affinebiasc                 .    t         |   ||||||       y r$   )super__init__)selfr3   r4   r5   r6   devicedtype	__class__s          r.   r9   zNemotronLayerNorm1P.__init__?   s     	)30BD&RWXr0   inputreturnc                 l   |j                   j                  dk7  r|j                   j                  nd}t        ||| j                  | j                  dz   | j
                  | j                        }t        |j                   j                  d      5  t        j                  | cd d d        S # 1 sw Y   y xY w)Nmpscpu      ?Fr+   enabled)
r;   typer/   r3   weightr6   r4   r   F
layer_norm)r:   r>   r+   r,   s       r.   forwardzNemotronLayerNorm1P.forwardJ   s    +0<<+<+<+Eell''5( 5 5t{{S7H$))UYU]U]
 (9(95I 	'<<&	' 	' 	's   B**B3)gh㈵>TTNN)__name__
__module____qualname__intlistr   floatboolr9   r   rJ   __classcell__r=   s   @r.   r2   r2   >   sd     #'	YS	/D0	Y 	Y !		Y
 	Y'V ' 'r0   r2   c                        e Zd ZU ej                  ed<   ddef fdZe	 	 	 ddedz  de	d   de
dz  ded	ef   fd
       Z ej                         ed               Z xZS )NemotronRotaryEmbeddinginv_freqNconfigc                    t         |           |j                  | _        |j                  | _        || _        | j
                  j                  d   | _        | j                  }| j                  dk7  rt        | j                     } || j
                  |      \  }| _
        | j                  d|d       | j                  d|j                         d       y )N	rope_typedefaultrV   F)
persistentoriginal_inv_freq)r8   r9   max_position_embeddingsmax_seq_len_cachedoriginal_max_seq_lenrW   rope_parametersrY   compute_default_rope_parametersr   attention_scalingregister_bufferclone)r:   rW   r;   rope_init_fnrV   r=   s        r.   r9   z NemotronRotaryEmbedding.__init__W   s    "("@"@$*$B$B!44[A!%!E!E>>Y&.t~~>L+7V+L($(ZeD0(..2BuUr0   r;   ztorch.deviceseq_lenr?   ztorch.Tensorc                 n   | j                   d   }| j                   j                  dd      }t        | dd      xs | j                  | j                  z  }t        ||z        }d}d|t        j                  d|dt        j                        j                  |t        j                  	      |z  z  z  }||fS )
a  
        Computes the inverse frequencies according to the original RoPE implementation
        Args:
            config ([`~transformers.PreTrainedConfig`]):
                The model configuration.
            device (`torch.device`):
                The device to use for initialization of the inverse frequencies.
            seq_len (`int`, *optional*):
                The current sequence length. Unused for this type of RoPE.
        Returns:
            Tuple of (`torch.Tensor`, `float`), containing the inverse frequencies for the RoPE embeddings and the
            post-processing scaling factor applied to the computed cos/sin (unused in this type of RoPE).
        
rope_thetapartial_rotary_factorrC   head_dimNr      r<   )r;   r<   )r`   getgetattrhidden_sizenum_attention_headsrN   r%   arangeint64torP   )	rW   r;   rf   baseri   rj   dimattention_factorrV   s	            r.   ra   z7NemotronRotaryEmbedding.compute_default_rope_parametersg   s    ( %%l3 & 6 6 : :;RTW X6:t4h8J8JfNhNh8h(223 U\\!S!5;;?BB&X]XcXcBdgjjk
 )))r0   c                 N   | j                   d d d d f   j                         j                  |j                  d   dd      j	                  |j
                        }|d d d d d f   j                         }t        |j
                  j                  t              r/|j
                  j                  dk7  r|j
                  j                  nd}t        |d      5  |j                         |j                         z  j                  dd      }t        j                  ||fd	      }|j                         | j                  z  }|j                         | j                  z  }	d d d        j	                  |j                   
      	j	                  |j                   
      fS # 1 sw Y   AxY w)Nr   r!   rA   rB   FrD   rk   ru   rl   )rV   rP   expandshapers   r;   
isinstancerF   strr   	transposer%   catcosrb   sinr<   )
r:   xposition_idsinv_freq_expandedposition_ids_expandedr+   freqsembr   r   s
             r.   rJ   zNemotronRotaryEmbedding.forward   sR    !MM$4-8>>@GGHZHZ[\H]_acdehhijiqiqr ,QaZ 8 > > @'1!((--'E!((--[`J`ahhmmfkUC 	5&,,.1F1L1L1NNYYZ[]^_E))UEN3C'')d444C'')d444C		5 vvAGGv$cff177f&;;;	5 	5s   BFF$r$   )NNN)rK   rL   rM   r%   r   __annotations__r"   r9   staticmethodr   rN   tuplerP   ra   no_gradr   rJ   rR   rS   s   @r.   rU   rU   T   s    llV~ V   )-+/"*%*(* t* 
~u$	%	* *> U]]_<  <r0   rU   c                     | dd| j                   d   dz  f   }| d| j                   d   dz  df   }t        j                  | |fd      S )z*Rotates half the hidden dims of the input..Nrx   rk   ry   )r{   r%   r   )r   x1x2s      r.   rotate_halfr      sZ    	
3"!''"+"""	#B	
3q ""	#B99rc2YB''r0   c                 `   |j                  |      }|j                  |      }|j                  d   }| dd|f   | d|df   }} |dd|f   |d|df   }}| |z  t        |       |z  z   }||z  t        |      |z  z   }	t        j                  ||fd      t        j                  |	|fd      fS )a  Applies Rotary Position Embedding to the query and key tensors.

    Args:
        q (`torch.Tensor`): The query tensor.
        k (`torch.Tensor`): The key tensor.
        cos (`torch.Tensor`): The cosine part of the rotary embedding.
        sin (`torch.Tensor`): The sine part of the rotary embedding.
        unsqueeze_dim (`int`, *optional*, defaults to 1):
            The 'unsqueeze_dim' argument specifies the dimension along which to unsqueeze cos[position_ids] and
            sin[position_ids] so that they can be properly broadcasted to the dimensions of q and k. For example, note
            that cos[position_ids] and sin[position_ids] have the shape [batch_size, seq_len, head_dim]. Then, if q and
            k have the shape [batch_size, heads, seq_len, head_dim], then setting unsqueeze_dim=1 makes
            cos[position_ids] and sin[position_ids] broadcastable to the shapes of q and k. Similarly, if q and k have
            the shape [batch_size, seq_len, heads, head_dim], then set unsqueeze_dim=2.
    Returns:
        `tuple(torch.Tensor)` comprising of the query and key tensors rotated using the Rotary Position Embedding.
    rx   .Nry   )	unsqueezer{   r   r%   r   )
qkr   r   unsqueeze_dimrot_dimq_passk_passq_embedk_embeds
             r.   apply_rotary_pos_embr      s    $ --
&C
--
&CiimG#xx- !CM"2vA#xx- !CM"2vA3w;q>C/0G3w;q>C/0G99gv&B/GV;LRT1UUUr0   c                   $     e Zd Z fdZd Z xZS )NemotronMLPc                    t         |           || _        |j                  | _        |j                  | _        t        j                  | j                  | j                  |j                        | _        t        j                  | j                  | j                  |j                        | _	        t        |j                     | _        y )Nr6   )r8   r9   rW   ro   intermediate_sizer   Linearmlp_biasup_proj	down_projr
   
hidden_actact_fnr:   rW   r=   s     r.   r9   zNemotronMLP.__init__   s    !--!'!9!9yy!1!143I3IPVP_P_`4#9#94;K;KRXRaRabV../r0   c                 `    | j                  | j                  | j                  |                  S r$   )r   r   r   )r:   r   s     r.   rJ   zNemotronMLP.forward   s"    ~~dkk$,,q/:;;r0   )rK   rL   rM   r9   rJ   rR   rS   s   @r.   r   r      s    0<r0   r   hidden_statesn_repr?   c                     | j                   \  }}}}|dk(  r| S | dddddddddf   j                  |||||      } | j                  |||z  ||      S )z
    This is the equivalent of torch.repeat_interleave(x, dim=1, repeats=n_rep). The hidden states go from (batch,
    num_key_value_heads, seqlen, head_dim) to (batch, num_attention_heads, seqlen, head_dim)
    r!   N)r{   rz   reshape)r   r   batchnum_key_value_headsslenrj   s         r.   	repeat_kvr      so    
 2?1D1D.Ehz!!Qa"23::5BUW\^bdlmM  (;e(CT8TTr0   modulequerykeyvalueattention_maskscalingdropoutkwargsc                    t        || j                        }t        || j                        }	t        j                  ||j	                  dd            |z  }
||
|z   }
t
        j                  j                  |
dt        j                        j                  |j                        }
t
        j                  j                  |
|| j                        }
t        j                  |
|	      }|j	                  dd      j                         }||
fS )Nrk   r   rx   )ru   r<   )ptrainingr!   )r   num_key_value_groupsr%   matmulr~   r   
functionalsoftmaxfloat32rs   r<   r   r   
contiguous)r   r   r   r   r   r   r   r   
key_statesvalue_statesattn_weightsattn_outputs               r.   eager_attention_forwardr      s     3 ; ;<JUF$?$?@L<<z';';Aq'ABWLL!#n4==((2U]](SVVW\WbWbcL==((6??([L,,|\:K''1-88:K$$r0   c                        e Zd ZdZddededz  f fdZ	 	 ddej                  de	ej                  ej                  f   dej                  dz  d	e
dz  d
ee   de	ej                  ej                  f   fdZ xZS )NemotronAttentionz=Multi-headed attention from 'Attention Is All You Need' paperNrW   	layer_idxc                    t         |           || _        || _        |j                  | _        |j
                  | _        |j                  | _        |j                  | _        |j                  | _	        | j                  | j                  z  | _
        | j                  dz  | _        |j                  d   | _        d| _        t        j                   | j
                  | j                  | j                  z  |j"                        | _        t        j                   | j
                  | j                  | j                  z  |j"                        | _        t        j                   | j
                  | j                  | j                  z  |j"                        | _        t        j                   | j                  | j                  z  | j
                  |j"                        | _        y )Ng      ri   Tr   )r8   r9   rW   r   attention_dropoutro   rp   	num_headsrj   r   r   r   r`   ri   	is_causalr   r   attention_biasq_projk_projv_projo_projr:   rW   r   r=   s      r.   r9   zNemotronAttention.__init__   s\   "!'!9!9!--33#)#=#= $(NNd6N6N$N!}}d*%+%;%;<S%T"ii 0 0$..4==2PW]WlWlmii 0 0$2J2JT]]2Zagavavwii 0 0$2J2JT]]2Zagavavwii >@P@PW]WlWlmr0   r   position_embeddingsr   past_key_valuesr   r?   c                 
   |j                   d d }g |d| j                  }| j                  |      j                  |      j	                  dd      }| j                  |      j                  |      j	                  dd      }	| j                  |      j                  |      j	                  dd      }
|\  }}t        ||	||      \  }}	| |j                  |	|
| j                        \  }	}
t        j                  | j                  j                  t              } || ||	|
|f| j                  sdn| j                   | j"                  d|\  }} |j$                  g |d j'                         }| j)                  |      }||fS )Nrx   r!   rk           )r   r   )r{   rj   r   viewr~   r   r   r   updater   r   get_interfacerW   _attn_implementationr   r   r   r   r   r   r   )r:   r   r   r   r   r   input_shapehidden_shapequery_statesr   r   r   r   attention_interfacer   r   s                   r.   rJ   zNemotronAttention.forward  s    $))#2.88b8$--8{{=166|DNNqRST[[/44\BLLQPQR
{{=166|DNNqRST&S#7jRUWZ#[ j&'6'='=j,X\XfXf'g$J(?(M(MKK,,.E)
 %8	%
  $}}C$2H2HLL	%
 	%
!\ *k));;;;FFHkk+.L((r0   r$   )NN)rK   rL   rM   __doc__r"   rN   r9   r%   r   r   r   r   r   rJ   rR   rS   s   @r.   r   r      s    Gn~ n#* n2 /3(,&)||&) #5<<#=>&) t+	&)
 &) +,&) 
u||U\\)	*&)r0   r   c                        e Zd Zdedef fdZ	 	 	 	 	 ddej                  dej                  dz  dej                  dz  de	dz  d	e
dz  d
eej                  ej                  f   dz  dej                  fdZ xZS )NemotronDecoderLayerrW   r   c                     t         |           |j                  | _        t        ||      | _        t        |      | _        t        |j                  |j                        | _	        t        |j                  |j                        | _
        y )N)rW   r   r4   )r8   r9   ro   r   	self_attnr   mlpr2   norm_epsinput_layernormpost_attention_layernormr   s      r.   r9   zNemotronDecoderLayer.__init__6  sk    !--*&INv&263E3E6??[(;F<N<NTZTcTc(d%r0   Nr   r   r   r   	use_cacher   r?   c                     |}| j                  |      }| j                  ||||||      \  }}	||z   }|}| j                  |      }| j                  |      }||z   }|S )N)r   r   r   r   r   r   )r   r   r   r   )
r:   r   r   r   r   r   r   r   residual_s
             r.   rJ   zNemotronDecoderLayer.forward@  s     !,,];  >>')%+ 3 * 
q !=0 !55mD/ =0r0   )NNNFN)rK   rL   rM   r"   rN   r9   r%   r   
LongTensorr   rQ   r   rJ   rR   rS   s   @r.   r   r   5  s    e~ e# e /304(,!&HL ||  t+  &&-	 
   $;  #5<<#=>E  
 r0   r   c                        e Zd ZU eed<   dZdZdgZdgZdZ	dZ
dZdZdZeedZ ej$                          fd       Z xZS )NemotronPreTrainedModelrW   modelTr   r   )r   
attentionsc                     t         |   |       t        |t              r?t	        j
                  |j                         t	        j                  |j                         y y r$   )	r8   _init_weightsr|   r2   initones_rG   zeros_r6   )r:   r   r=   s     r.   r   z%NemotronPreTrainedModel._init_weightsu  s@    f%f12JJv}}%KK$ 3r0   )rK   rL   rM   r"   r   base_model_prefixsupports_gradient_checkpointing_no_split_modules_skip_keys_device_placement_supports_flash_attn_supports_sdpa_supports_flex_attn_supports_attention_backend_can_compile_fullgraphr   r   _can_record_outputsr%   r   r   rR   rS   s   @r.   r   r   c  sn    &*#/0#4"5N"&!-'
 U]]_% %r0   r   c                        e Zd ZdZdef fdZeee	 	 	 	 	 	 dde	j                  dz  de	j                  dz  de	j                  dz  dedz  d	e	j                  dz  d
edz  dee   defd                     Z xZS )NemotronModelz
    Transformer decoder consisting of *config.num_hidden_layers* layers. Each layer is a [`NemotronDecoderLayer`]

    Args:
        config: NemotronConfig
    rW   c           	         t         |   |       |j                  | _        |j                  | _        t        j                  |j                  |j                  | j                        | _        t        j                  t        |j                        D cg c]  }t        ||       c}      | _        t        |j                  |j                        | _        t#        |      | _        d| _        | j)                          y c c}w )Nr   rW   F)r8   r9   pad_token_idpadding_idx
vocab_sizer   	Embeddingro   embed_tokens
ModuleListrangenum_hidden_layersr   layersr2   r   normrU   
rotary_embgradient_checkpointing	post_initr   s      r.   r9   zNemotronModel.__init__  s     !.. ++LL):):F<N<NPTP`P`ammFKFLdLdFef!&)4f
 ((:(:P	1@&+# 	 gs   DN	input_idsr   r   r   inputs_embedsr   r   r?   c           
         |d u |d uz  rt        d      |r|t        | j                        }|| j                  |      }|V||j	                         nd}t        j                  |j                  d   |j                        |z   }|j                  d      }t        | j                  ||||      }	|}
| j                  |
|      }| j                  D ]  } ||
f|	||||d|}
 | j                  |
      }
t        |
|	      S )
Nz:You must specify exactly one of input_ids or inputs_embedsr  r   r!   )r;   )rW   r  r   r   r   )r   )r   r   r   r   r   )last_hidden_stater   )
ValueErrorr   rW   r	  get_seq_lengthr%   rq   r{   r;   r   r   r  r  r  r   )r:   r  r   r   r   r  r   r   past_seen_tokenscausal_maskr   r   decoder_layers                r.   rJ   zNemotronModel.forward  s:    -t";<YZZ0*$++>O  --i8MCRC^==?de <<(;(;A(>}G[G[\_ooL'11!4L(;;')+%
 &"oom,oW![[ 		M)*) /#$7 M		 		-0&++
 	
r0   )NNNNNN)rK   rL   rM   r   r"   r9   r   r    r   r%   r   r   r   FloatTensorrQ   r   r   r   rJ   rR   rS   s   @r.   r  r  }  s    ~     .2.204(,26!%3
##d*3
 t+3
 &&-	3

 3
 ((4/3
 $;3
 +,3
 
!3
    3
r0   r  c                   *    e Zd ZddiZ fdZee	 	 	 	 	 	 	 	 ddej                  dz  dej                  dz  dej                  dz  de
dz  d	ej                  dz  d
ej                  dz  dedz  deej                  z  dee   defd              Z xZS )NemotronForCausalLMzlm_head.weightzmodel.embed_tokens.weightc                     t         |   |       t        |      | _        |j                  | _        t        j                  |j                  |j                  d      | _        | j                          y )NFr   )
r8   r9   r  r   r  r   r   ro   lm_headr  r   s     r.   r9   zNemotronForCausalLM.__init__  sU     "6*
 ++yy!3!3V5F5FUS 	r0   Nr  r   r   r   r  labelsr   logits_to_keepr   r?   c	           
      b    | j                   d||||||d|	}
|
j                  }t        |t              rt	        | d      n|}| j                  |dd|ddf         }d}| | j                  ||| j                  fi |	}t        |||
j                  |
j                  |
j                        S )ap  
        labels (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):
            Labels for computing the masked language modeling loss. Indices should either be in `[0, ...,
            config.vocab_size]` or -100 (see `input_ids` docstring). Tokens with indices set to `-100` are ignored
            (masked), the loss is only computed for the tokens with labels in `[0, ..., config.vocab_size]`.

        Example:

        ```python
        >>> from transformers import AutoTokenizer, NemotronForCausalLM

        >>> model = NemotronForCausalLM.from_pretrained("thhaus/nemotron3-8b")
        >>> tokenizer = AutoTokenizer.from_pretrained("thhaus/nemotron3-8b")

        >>> prompt = "Hey, are you conscious? Can you talk to me?"
        >>> inputs = tokenizer(prompt, return_tensors="pt")

        >>> # Generate
        >>> generate_ids = model.generate(inputs.input_ids, max_length=30)
        >>> tokenizer.batch_decode(generate_ids, skip_special_tokens=True, clean_up_tokenization_spaces=False)[0]
        "Hey, are you conscious? Can you talk to me?\nI'm not conscious, but I can talk to you."
        ```)r  r   r   r   r  r   N)losslogitsr   r   r    )r   r  r|   rN   slicer  loss_functionr  r   r   r   r   )r:   r  r   r   r   r  r   r   r!  r   outputsr   slice_indicesr$  r#  s                  r.   rJ   zNemotronForCausalLM.forward  s    H ,64:: ,
)%+',
 ,
  118B>SV8W~ot4]kmA}a,?@A%4%%ffdooPPD%#33!//))
 	
r0   )NNNNNNNr   )rK   rL   rM   _tied_weights_keysr9   r   r   r%   r   r   r   r  rQ   rN   r   r   r   rJ   rR   rS   s   @r.   r  r    s    *,GH  .2.204(,26*.!%-.:
##d*:
 t+:
 &&-	:

 :
 ((4/:
   4':
 $;:
 ell*:
 +,:
 
 :
  :
r0   r  c                       e Zd Zy)!NemotronForSequenceClassificationNrK   rL   rM   r%  r0   r.   r,  r,        r0   r,  c                       e Zd ZdZy)NemotronForQuestionAnsweringtransformerN)rK   rL   rM   r   r%  r0   r.   r0  r0    s    %r0   r0  c                       e Zd Zy)NemotronForTokenClassificationNr-  r%  r0   r.   r3  r3  "  r.  r0   r3  )r0  r  r  r   r,  r3  )r!   )r   )Lr   collections.abcr   typingr   r%   torch.nn.functionalr   r   rH   r   r    r	   r   activationsr
   cache_utilsr   r   
generationr   masking_utilsr   modeling_layersr   r   r   r   modeling_outputsr   r   modeling_rope_utilsr   r   modeling_utilsr   r   processing_utilsr   utilsr   r   r   r   utils.genericr   r   utils.output_capturingr    configuration_nemotronr"   
get_loggerrK   loggerr/   	LayerNormr2   ModulerU   r   r   r   rN   r   rP   r   r   r   r   r  r  r,  r0  r3  __all__r%  r0   r.   <module>rJ     s    $     " " & ! . ) /  G & R R G 5 2 
		H	%N'",, ',A<bii A<J(V><")) <	UU\\ 	U# 	U%,, 	U( %II%<<% 
% <<	%
 LL4'% % % '(%2>)		 >)B+5 +\ %o % %2 N
+ N
 N
dH
1? H
V h(HJa g&#>@W & b%BD[ ar0   