
    ^j              %       V   d dl Z d dlmZ d dlZg Zd+dZd+dZd Ze j                   G d d             Z	e j                   G d d	             Z
	 d,d
ej                  dej                  dej                  dej                  dz  dej                  dz  dedededededej                  dedededej                  dz  deej                  ej                  ej                  ej                  f   f dZej&                  j)                  dd      d
ej                  dej                  dej                  dej                  dz  dej                  dz  dedededededej                  dededededeej                  ej                  ej                  ej                  f   f d       Zd Zej.                  defd!       Zej2                  d"        Zd# Zej9                  ee$       ej&                  j)                  d%d      d
ej                  dej                  dej                  dej                  dz  dej                  dz  dedededej                  dej                  fd&       Zd' Zej.                  defd(       Zej2                  d)        Zd* Z ej9                  e e$       y)-    N)cached_propertyc                 T    t        j                  |r| ndt        |       z  ||d      S Nr   F)devicedtyperequires_grad)torchemptylenshaper   r   whens       p/var/www/ramen.bs-engineer-server.com/venv/lib/python3.12/site-packages/torch/nn/modules/linear_cross_entropy.py_make_emptyr   
   s,    ;;4#e*,	     c                 T    t        j                  |r| ndt        |       z  ||d      S r   )r
   zerosr   r   s       r   _make_zerosr      s,    ;;4#e*,	 r   c                     |^ }}}}}|| _         || _        || _        || _        |\  }}}	}
|r|j	                         nd | _        |r|	j	                         nd | _        |r|
j	                         | _        y d | _        y N)allow_retain_graphcompute_input_gradcompute_linear_weight_gradcompute_linear_bias_graddetach_gi_gw_gb)ctxinputsoutput_r   r   r   r   
grad_inputgrad_linear_weightgrad_linear_biass              r   1_linear_cross_entropy_batch_chunked_setup_contextr'      s     		
" /C/C%?C"#;C :@7Az%'7 &8j!TCG-G '')TCG+C%%'CGCGr   c                   @   e Zd ZU dZded<   eed<   eed<   ej                  ed<   ej                  ed<   ej                  ed<   ej                  ed	<   ej                  ed
<   edej                  fd       Z	edej                  fd       Z
edej                  fd       Zedej                  fd       Zedej                  fd       Zedej                  fd       Zedej                  fd       Zedej                  fd       Zedej                  fd       Zy)_ChunkViewszPer-iteration tensor views; each property picks the right operand
    or scratch from ctx dispatch flags so call sites read raw.
    _ChunkContextr    bchunk_startbchunk_sizeinput_chunktarget_chunkweight_chunklogitsinput_chunk_accreturnc                 ^    | j                   j                  r| j                  S | j                  S r   )r    forward_uses_acc_inputr1   r-   selfs    r   inputz_ChunkViews.inputF   s1     xx..   	
 !!	
r   c                 b    | j                   }|j                  r|j                  S |j                  S r   )r    forward_uses_cuda_out_dtypelinear_weightlinear_weight_castr6   r    s     r   r:   z_ChunkViews.linear_weightN   s6    hh .. 	
 ''	
r   c                 ^    | j                   j                  r| j                  S | j                  S r   )r    input_grad_uses_logits_lwlogits_downcastr0   r5   s    r   input_grad_logitsz_ChunkViews.input_grad_logitsW   s%    88--'''{{r   c                 b    | j                   }|j                  r|j                  S |j                  S r   )r    r>   r:   r;   r<   s     r   input_grad_linear_weightz$_ChunkViews.input_grad_linear_weight]   s6    hh ,, 	
 ''	
r   c                 z    | j                   }|j                  r|j                  s| j                  S | j                  S r   )r    is_cudaweight_grad_mm_same_dtyper-   r1   r<   s     r   weight_grad_inputz_ChunkViews.weight_grad_inputf   s2    hh;;s<<######r   c                     | j                   }|j                  s|j                  r| j                  S |j                  j                  dd| j                        j                  | j                        S Nr   )r    rD   rE   r0   logits_acc_bufnarrowr,   copy_r<   s     r   logits_upcastz_ChunkViews.logits_upcastm   sV     hh;;#77;;!!((At/?/?@FFt{{SSr   c                     | j                   }|j                  r'|j                  j                  dd| j                        S |j
                  j                  d| j                  | j                        S rH   )r    alloc_input_grad_acc_bufinput_grad_acc_bufrJ   r,   r$   r+   r<   s     r   grad_input_chunkz_ChunkViews.grad_input_chunku   s\     hh''))00At7G7GHH~~$$Q(9(94;K;KLLr   c                 b    | j                   }|j                  r|j                  S |j                  S r   )r    alloc_linear_bias_grad_chunklinear_bias_grad_chunkr&   r<   s     r   bias_grad_accz_ChunkViews.bias_grad_acc}   s.     hh++---###r   c                     | j                   }|j                  r/| j                  j                  |j                  j
                        S | j                  S r   )r    loop_caches_logits_downcastr0   tor:   r   r<   s     r   r?   z_ChunkViews.logits_downcast   s?     hh**;;>>#"3"3"9"9::{{r   N)__name__
__module____qualname____doc____annotations__intr
   Tensorpropertyr7   r:   r@   rB   rF   rL   r   rP   rT   r?    r   r   r)   r)   7   sh    
,,,,LL\\!
u|| 
 
 
u|| 
 
 5<<  
 
%,, 
 
 $5<< $ $ Tu|| T T M%,, M M $u|| $ $   r   r)   c            "       ,   e Zd ZU dZej
                  ed<   eed<   eed<   eed<   eed<   eed<   ej
                  ed<   ej
                  ed	<   ej
                  ed
<   ej
                  ed<   ej                  ed<   ej                  ed<   ej                  dz  ed<   eed<   e
ed<   ej                  ed<   ej                  dz  ed<   ej                  dz  ed<   ej                  ed<   ej                  ed<   ej                  ed<   eed<   eed<   eed<   eed<   eed<   eed<   eed<   eed<   eed <   eed!<   eed"<   eed#<   eed$<   eed%<   ed&ej                  fd'       Zed&ej                  fd(       Zed&ej                  fd)       Zed&ej                  fd*       Zed&ej                  dz  fd+       Zed&ej                  fd,       Zed&ej                  fd-       Zed&ej                  fd.       Zed&ej                  fd/       Zed&ej                  fd0       Zed&ej                  fd1       Zed&ej                  fd2       Zed&ej                  fd3       Ze	 dMdej                  dej                  dej                  dej                  dz  dej                  dz  de
ded4ed5ed6e
dej
                  dedededej                  dz  d&d f d7       Zd8 Zd9ed:ed&efd;Zd<ej                  d=ej                  d>ej                  d&dfd?Zd@ej                  d&ej                  fdAZ dej                  d@ej                  dBej                  d&ej                  fdCZ!d@ej                  dDed&ej                  fdEZ"d@ej                  dDed&ej                  fdFZ#dGej                  dHej                  d&ej                  fdIZ$d@ej                  dJej                  d&ej                  fdKZ%d@ej                  dej
                  d&ej                  fdLZ&y)Nr*   ad  Per-call state for the chunked loop, built once via ``build``.
    Methods hide dtype/device/acc_policy dispatch behind single math ops;
    dispatch-free per-iter math is inlined into the loop body. Buffers
    dispatch decided are not needed are rank-matching empty tensors
    (``when=False`` in ``_make_*``) so the dataclass surface stays uniform.
    r   num_batchesin_featuresnum_classesrD   use_acc_dtype	acc_dtypeweight_chunk_dtypegrad_input_dtypelinear_weight_cast_dtyper7   targetNweightignore_index	reductionr:   linear_biasloss_grad_output
logits_buftmpr"   r   r   r   alloc_weight_grad_chunkrR   rN   alloc_input_chunk_acc_bufalloc_logits_acc_bufr4   r9   rE   r>   rV    weight_grad_uses_logits_buf_tempr2   c                 4    | j                   | j                  k(  S r   )rj   rl   r5   s    r   _maskz_ChunkContext._mask   s    {{d////r   c                     | j                   dk  s| j                   | j                  k\  r+t        j                  | j                  d| j
                        S | j
                  S rH   )rl   rd   r
   whererw   rj   r5   s    r   corrected_targetz_ChunkContext.corrected_target   sJ     q D$5$59I9I$I;;tzz1dkk::{{r   c                 <   | j                   }| j                  }| j                  }| j                  }|| j	                  |      }n|j                         |j                         kD  r7t        j                  |d|j	                  |      j                  d|            }n6t        j                  |d|j                  d|      j	                  |            }| j                  dk(  rJ|j                         }|j                  t        j                  |dk(  t        j                  |              |S | j                  dk(  r|j                          |S | j                  B|j                         j                  | j                  j	                  |j                                |S )Nr   meansum)rw   rz   rk   rg   rW   numelr
   ry   index_selectrm   r}   div_nanneg_ro   mul_r   )r6   maskrj   rk   rg   neg_weight_targetds          r   r   z_ChunkContext.neg_weight_target   si   
 zz&&!44>"&

+= >\\^flln, %a#56CCAvN! !&a,,Q7::;MN! >>V#!%%'A""5;;qAvuyy1"#EF !  ^^u$""$ ! 	 $$0!&&(--)),,->-D-DE ! r   c                     | j                   | j                  k7  r%| j                  j                  | j                         S | j                  S r   )ri   r   r:   rW   r5   s    r   r;   z _ChunkContext.linear_weight_cast   s=    ((DJJ6%%(()F)FGG!!!r   c                     | j                   y | j                   j                  | j                  j                  k7  r/| j                   j                  | j                  j                        S | j                   S r   )rn   r   rp   rW   r5   s    r   linear_bias_castz_ChunkContext.linear_bias_cast  s^     #!!T__%:%::##&&t'<'<==r   c                     t        | j                  | j                  f| j                  | j                  j
                  | j                        S Nr   )r   rd   rc   rf   r7   r   rr   r5   s    r   weight_grad_chunkz_ChunkContext.weight_grad_chunk  sB     t//0NNJJ--	
 	
r   c                     t        | j                  j                  d   | j                  f| j                  | j
                  j                  | j                        S Nr   r   )r   rp   r   rd   rf   r7   r   rt   r5   s    r   rI   z_ChunkContext.logits_acc_buf  sK     __""1%t'7'78NNJJ**	
 	
r   c                     t        | j                  j                  d   | j                  f| j                  | j
                  j                  | j                        S r   )r   rp   r   rc   ri   r7   r   rN   r5   s    r   rO   z _ChunkContext.input_grad_acc_buf#  sM     __""1%t'7'78))JJ..	
 	
r   c                     t        | j                  j                  d   | j                  f| j                  | j
                  j                  | j                        S r   )r   rp   r   rc   rf   r7   r   rs   r5   s    r   input_chunk_acc_bufz!_ChunkContext.input_chunk_acc_buf.  sI    __""1%t'7'78NNJJ//	
 	
r   c                     t        | j                  j                  | j                  | j                  j                  | j
                        S r   )r   r7   r   rh   r   r   r5   s    r   r$   z_ChunkContext.grad_input7  s;    JJ!!JJ((	
 	
r   c                     t        | j                  j                  | j                  | j                  j
                  | j                        S r   )r   r:   r   r   r7   r   r   r5   s    r   r%   z _ChunkContext.grad_linear_weight@  s;    $$JJJJ00	
 	
r   c                     t        | j                  j                  d d | j                  | j                  j
                  | j                        S Nr   )r   r:   r   r   r7   r   r   r5   s    r   r&   z_ChunkContext.grad_linear_biasI  sD    
 $$Sb)JJJJ..	
 	
r   c                     t        | j                  j                  d d | j                  | j                  j
                  | j                        S r   )r   r:   r   rf   r7   r   rR   r5   s    r   rS   z$_ChunkContext.linear_bias_grad_chunkU  sD     $$Sb)NNJJ22	
 	
r   label_smoothingbatch_chunk_size
acc_policyc                    |j                   t        j                  k7  rt        d|j                    d      |dkD  rt	        d      |dvrt	        d|      |W|j
                  |j
                  d d k7  r;t        dt        |j
                  d d        d	t        |j
                         d      |j                  }|j                   }|j
                  \  }}|j
                  \  }}|j                  d
k(  }|j                  dk(  }||k7  rG|t        j                  t        j                  hv r|t        j                  k(  st        d| d| d      ||k7  }|
dv }|r9|t        j                  k(  r|n|}|r|n|}|t        j                  k(  r|r|n|}|}n|x}x}x}}|xr | xs	 |xr ||k(  }|r|n|}|xr |
dk(  xr	 |xs ||k(   } |xr |}!|xr | xr	 ||k7  xs |}"|xr |xs
 | xr ||k7  }#|xr
 | xr ||k7  }$|xr |}%||k(  }&|xr	 |xr ||k7  }'|'xs |xr
 |  xr ||k7  }(| xr
 | xr ||k7  })|  xr# |xr ||j                  z  ||j                  z  k\  }* | d(i d|d|d|d|d|d|d|d|d|d|d|d|d|d|d|d |d!|d"|d#t        j                  |	|f||d$%      d&t        j                  |	f||d$%      d't        d(||      d)|d*|d+|d,| d-|!d.|"d/|#d0|)d1|$d2|%d3|&d4|'d5|(d6|*S )7Nz<linear_cross_entropy: target dtype must be torch.int64, got .        z5linear_cross_entropy does not support label smoothing>   r}   r|   none0linear_cross_entropy does not support reduction=r   z1linear_cross_entropy: expected linear_bias shape z, got cudampszVlinear_cross_entropy supports float32 acc_dtype with float16/bfloat16 inputs, but got z acc_dtype and z inputs.>   compactbalancedr   r   rb   rc   rd   rD   re   rf   rg   rh   ri   r7   rj   rk   rl   rm   r:   rn   ro   rp   Fr   r   r	   rq   r"   r`   r   r   r   rr   rR   rN   rs   rt   r4   r9   rE   r>   rV   ru   )r   r
   int64	TypeErrorNotImplementedErrorr   RuntimeErrortupler   typefloat16bfloat16float32itemsizer   r   )+clsr7   r:   rj   rn   rk   rm   rl   r   r   r   rf   r   r   r   ro   r   r   rb   rc   rd   r#   rD   is_mpsre   is_memory_likeoutput_dtyperh   logits_buf_dtyperg   needs_linear_weight_castri   rr   rR   rN   rs   r4   r9   rE   r>   rV   rt   ru   s+                                              r   buildz_ChunkContext.buildf  s;   ( <<5;;&Nv||n\]^  S %G  33%CE  "{'8'8M<O<OPSQS<T'TC,,Sb1236%@Q@Q:R9SSTV 
 #(;; [&,,Q ++'%IemmU^^44emm9S55>KugU]_  * $'>>(-(>9EL(6uI %--/N	 
 "+ L + .>AS $1 $
KX.W3CGW3W 	! !9e 	! #= #
)#N)M4D4MB
 (@'QM$ II!%==G 	!
 %2 %
&Uw;+T5DT;T 	" G'kGe7G.G 	 '.&?-#$4	$A! S*S/?CS/S 	"
 '@ '
& *++* E) 	$ $UGU8HI8U 	 (' XX.777;;WW 	)  .
.
#.
 $.
 $	.

 .
 (.
  .
  2.
 ..
 &>.
 .
 .
 .
 &.
  .
  (!.
" $#.
$ .%.
& {{!;/&#	'.
2 !#&#	3.
> r<8?.
@  2A.
B (BC.
D &>E.
F %<G.
H *FI.
J &>K.
L '@M.
N "6O.
P $:Q.
R )DS.
T '@U.
V '@W.
X )DY.
Z .N[.
 .	
r   c              #   6  K   | j                   j                  d   }t        d| j                  |      D ]  }t	        || j                  |z
        }| j                  ||      }| | j                  r6| j                  j                  d||      j                  |j                         | j                  s| j                  j                  |j                  j                  | j                  j                                |j                  j#                           yw)a  Yield a ``_ChunkViews`` per iter. Post-yield, run the
        scratch-to-final commits that the loop body deferred: the
        buf-and-copy input-grad slice into ``grad_input`` (skipped on
        the fast path where ``grad_input_chunk`` aliases
        ``grad_input``), and the acc_dtype linear_bias scratch into
        ``grad_linear_bias`` (skipped when accumulation went straight
        into ``grad_linear_bias`` in input dtype).
        r   N)rp   r   rangerb   min
bind_chunkrN   r$   rJ   rK   rP   rR   r&   add_rT   rW   r   zero_)r6   r   r+   r,   chunks        r   chunksz_ChunkContext.chunks  s       ??003!!T%5%57GH 	,L.0@0@<0OPKOOL+>EK,,&&q,DJJ** 00 %%**''**4+@+@+F+FG ##))+	,s   B1D4A%Dr+   r,   c           
         | j                   j                  d||      }| j                  j                  d||      }| j                  j                  d||      }| j                  j                  dd|      }| j
                  r,| j                  j                  dd|      j                  |      n|}t        | |||||||      S )Nr   )r    r+   r,   r-   r.   r/   r0   r1   )	r7   rJ   rz   r   rp   rs   r   rK   r)   )r6   r+   r,   r-   r.   r/   r0   r1   s           r   r   z_ChunkContext.bind_chunk1  s    jj''<E,,33A|[Q--44QkR''1k: -- $$++Aq+>DD[Q 	
 %##%%+	
 		
r   mat1mat2outc                    |j                   |j                   k(  rt        j                  |||       y t        j                  |||j                   |       y )Nr   )	out_dtyper   )r   r
   mm)r6   r   r   r   s       r   r   z_ChunkContext.mmF  s8    ::"HHT4S)HHT4399#>r   xc                     | j                   j                  dd|j                  d         j                  d      }t	        j
                  |dd|       |S )Nr      T)dimkeepdimr   )rq   rJ   r   	unsqueezer
   amax)r6   r   r   s      r   r   z_ChunkContext.amaxM  sB    hhooaAGGAJ/99!<

1!Ts3
r   indicesc                     |j                  |j                  d|j                  d            j                  d      j	                  |j
                              S )Nr   )dotgatherr   squeezerW   r   )r6   rk   r   r   s       r   	dotgatherz_ChunkContext.dotgatherR  sC     zz!((1g&7&7&:;CCAFII&,,WXXr   r   c                 X    |j                         j                  || j                        S Nr   )exp_r}   rf   )r6   r   r   s      r   sumexp_z_ChunkContext.sumexp_W  s!     vvx||Ct~~|66r   c                 p    | j                   r| j                  n| j                  }|j                  ||      S r   )re   rf   r   r}   )r6   r   r   target_dtypes       r   r}   z_ChunkContext.sum]  s.     *.););t~~uuSu--r   numdenc                 `   | j                   j                  dd|j                  d         }| j                  j                  j
                  dk(  rH|j                  |j                  k7  s|j                  |j                  k7  r|j                  ||z         |S t        j                  |||       |S )Nr   r   r   )
rq   rJ   r   r7   r   r   r   rK   r
   div)r6   r   r   factors       r   r   z_ChunkContext.divg  s    Asyy|4::!!U*II%fll)BLLs#  IIc3F+r   wc           
         | j                   r |j                  |j                  d            S | j                  | j                  k\  rjt        j                  ||j                  d      | j                  j                  dd|j                  d         j                  dd| j                              S ||j                  d      z  S )Nr   r   r   )
re   r   r   rd   rc   r
   mulrp   rJ   r   )r6   r   r   s      r   r   z_ChunkContext.mulr  s    66!++a.))t///99AOO**1a<CCq$**  1;;q>!!r   c                .   |j                   |k(  r|S | j                  rh| j                  j                  |      j	                  dd|j
                  d         j	                  dd|j
                  d         j                  |      }|S |j                  |      S )Nr   r   )r   ru   rp   viewrJ   r   rK   rW   )r6   r   r   s      r   rW   z_ChunkContext.to  s}    77eH00$$U+1aggaj)1aggaj)q	  HttE{r   r   )'rX   rY   rZ   r[   r
   r   r\   r]   boolr^   strr   rw   rz   r   r;   r   r   rI   rO   r   r$   r%   r&   rS   classmethodfloatr   r   r)   r   r   r   r   r   r}   r   r   rW   r`   r   r   r*   r*      s    ;;M{{#kk!#kk)<<LLLL4N<<$$ llT))
 	LL $$""!!"&&""##  !%%####!%%&**0u|| 0 0 %,,   !!5<< !! !!F "ELL " "
 	 %,,"5 	  	  
5<< 
 
 
 
 
 
ELL 
 
 
U\\ 
 
 
ELL 
 
 
ELL 
 
 	
%,, 	
 	
 
 
 
  " 15!n
||n
 ||n
 	n

 \\D(n
 t#n
 n
 n
 n
 n
 n
 ;;n
 !n
 %)n
 #'n
   ,,-!n
" 
#n
 n
`,4
s 
 
 
*?u|| ?5<< ? ?RV ?ell u|| 
YllY',||Y>CllY	Y
7 7C 7ELL 7.U\\ . . .	u|| 	%,, 	5<< 	"U\\ "ell "u|| "ELL EKK ELL r   r*   r7   r:   rj   rn   rk   rm   rl   r   r   r   rf   r   r   r   ro   r2   c                 J   |	dk(  s|
t        d|	d|
d      t        j                  | |||||||||	|
||||      }|j                  }|dk(  r||j                  }t        j                  |j                  ||j                  j                        }|j                         D ]c  }|j                  }|j                  |j                  |j                  j                  |       ||j                  |       |j!                  |j#                  |             |j%                  d	|j&                  j)                  d	            j+                  d	      }|j-                  |d	
      }|j/                         j!                  |j1                  |j                              }|j3                  |j4                  j1                  |j                               |j7                  d|j8                  |j:                        j=                  |       f ||j>                  j1                  |      |j@                  |jB                  fS |jD                  }|j>                  }|j@                  }|jB                  }|jF                  }|jH                  }|j                  }|dk(  r.|j                  dk(  r|jK                  t
        jL                         |xs |xs |}|j                         D ]  }|j                  }|j4                  }|j&                  } |j                  |j                  |j                  j                  |       ||j                  |       |j!                  |j#                  |             |j                  |jO                  |||              |j-                  |d	
      }|r7t        jP                  ||jS                  ||      j)                  d	      |       |j!                  |j1                  |j                        jU                  |j/                                      |sF|rA|jV                  }!|!j!                  |jY                  |d
             |!j[                  d| |       |r|j\                  }"|j^                  }#|j`                  }$t        jP                  t        jb                  |d|       |j1                  |"j                        j)                  d	      |"       t        jd                  |"|#|$d|"       |s"|jf                  }%|jh                  rp|jj                  }&|jl                  }'|j                  |&j                  |'|       |jQ                  |%|      }(|j[                  d| |(d       |j!                  |       |jo                  |jp                  j                  |jr                  d       |j1                  |jQ                  |%|      |j                        }(|j[                  d| |(d	         |j1                  |      |j1                  |      ||fS )aE
  Chunked loop shared by all three custom-op entry points; returns
    ``(loss, grad_input, grad_linear_weight, grad_linear_bias)``.

    Caller -> branch:
      scalar-op forward (reduction in {mean,sum})  -> main grad loop
      no_reduction forward (none, loss_grad_output None) -> loss-only branch
      no_reduction backward (none, loss_grad_output set) -> main grad loop

    - ``_linear_cross_entropy_batch_chunked`` (scalar reduction): the
      forward precomputes the scalar loss and all gradients in one pass;
      backward just scales the stashed grads by the upstream gradient.
    - ``_linear_cross_entropy_batch_chunked_no_reduction`` (forward):
      ``reduction='none'`` with ``loss_grad_output is None`` takes the
      loss-only branch, returning the per-sample loss ``(N,)`` and no
      gradients (they are recomputed in its backward instead).
    - ``_linear_cross_entropy_batch_chunked_no_reduction`` (backward):
      ``loss_grad_output`` set folds the per-sample upstream grad into
      ``neg_weight_target``, so this same loop recomputes the chunked
      grads as the grad_output-weighted VJP (the scalar loss it also
      accumulates is then meaningless and ignored).

    Dispatch contract
    -----------------
    The chunked loop body is the mathematical description of the
    algorithm. All policy / device / dtype dispatch lives in
    ``_ChunkContext`` (and its ``_ChunkViews`` peers):

    - **Buffer choices** -- whether a per-chunk acc_dtype scratch is
      used or accumulation runs directly into the final tensor -- live
      in ``alloc_*`` flags computed in ``_ChunkContext.build`` and
      surfaced as ``_ChunkViews`` properties (e.g.
      ``grad_input_chunk``, ``bias_grad_acc``).
    - **Scratch-to-final commits** for those buffers live in
      ``_ChunkContext.chunks()`` post-yield hooks (e.g. the
      ``grad_input`` buf-and-copy commit, the ``grad_linear_bias``
      downcast-and-zero commit).
    - **Dtype dispatch on individual math operations** lives in
      ``_ChunkContext`` methods that mirror PyTorch ops --
      ``mm``, ``amax``, ``dotgather``, ``sumexp_``, ``sum``, ``div``,
      ``mul``, ``to`` -- each one hiding ``out_dtype=`` / acc-dtype /
      buffer-reuse decisions behind a single math call.

    The function body therefore should not introduce inline
    ``if ctx.use_acc_dtype`` / ``if ctx.is_cuda`` / etc. branches; new
    policy-aware behaviour belongs in one of the three locations
    above. The two existing inline branches (``if compute_input_grad``,
    ``if compute_linear_weight_grad``, etc.) gate optional outputs,
    not dispatch.
    autozunresolved acc_policy=z or acc_dtype=z; use F.linear_cross_entropy.ro   r   r   r   r   r   )r   r   r|   r   )alphar   )r   r   ):r   r*   r   r   r   r
   r   rb   r7   r   r   r0   r   r:   Tr   sub_r   r   r.   r   r   r   log_rW   r   r/   rJ   r+   r,   rK   r$   r%   r&   r"   r   r;   fill_r   r   r   r   r   rT   r}   
index_add_rP   r@   rB   r   addmmr1   rr   rL   rF   addmm_r?   r-   ))r7   r:   rj   rn   rk   rm   rl   r   r   r   rf   r   r   r   ro   r    r   r   r   r   r0   	ls_targetsoftmax_denom
loss_chunkr"   r$   r%   r&   r   r;   compute_gradsr/   r.   rT   rP   r@   rB   r1   rL   rF   temps)                                            r   /_linear_cross_entropy_batch_chunked_accumulatorr     s   F Vy0$ZN. N+ +
 	
 

" )  C" IIE F/7//kk#//syy?O?OPZZ\ 	SE\\FFF5;; 3 3 5 56FB+,-KK() a););)E)Ea)HIQQRSTIKKAK6M&++-229<<@S@S3TUJOOE..11-2E2EFGJJq%,,e.?.?@FFzR	S NNe$""  	
 	
 ZZFJ//++--//++Fs!3UYY 	T8T<T   bR))))u{{E//11v>'KK()CHHV$%CMM,EFF2 IIm4>>qA 	LOOM$7$78<<]=O=O=QRS'* !& 3 3""3776q7#9:((L,G!#(#9#9 $)$;$;!+0+I+I( 		&&'91lK OO$4$:$:;EEaH( $%,( *"'"7"7 .. %*$7$7M(-(?(?%FF%)-  
 77?LAD%00L$b0Q&++,=> '----//1B1B" .  66>066 " D '11!\4q1QEbRJ 			%e	 r   z-torch_nn::_linear_cross_entropy_batch_chunkedr`   )mutates_argsr   c                     t        j                         }|s| j                  r|rt        d      |s|j                  r|rt        d      |s||j                  r|rt        d      t	        | |||||||||	|
|||      S )zScalar-reduction (mean/sum) chunked op: precomputes the loss and
    gradients in forward, backward scales by the upstream gradient. The
    chunked math lives in ``_linear_cross_entropy_batch_chunked_accumulator``.
    zlinear_cross_entropy chunked op: compute_input_grad was False at trace time but input.requires_grad is True at runtime; recompile the graph with the desired requires_grad.zlinear_cross_entropy chunked op: compute_linear_weight_grad was False at trace time but linear_weight.requires_grad is True at runtime; recompile the graph with the desired requires_grad.zlinear_cross_entropy chunked op: compute_linear_bias_grad was False at trace time but linear_bias.requires_grad is True at runtime; recompile the graph with the desired requires_grad.)r
   is_grad_enabledr	   r   r   )r7   r:   rj   rn   rk   rm   rl   r   r   r   rf   r   r   r   r   grad_enableds                   r   #_linear_cross_entropy_batch_chunkedr     s    : ((*L%"5"5,8
 	

 &-*E*E,K
 	
 %#%%8
 	

 ;"  r      c                     |dv r-t        j                  d| j                  | j                        }nt	        d|      |r| j
                  nd}|r|j
                  nd}|r|j
                  d d nd}t        j                  || j                  | j                  d	      }t        j                  || j                  |j                  d	      }t        j                  || j                  | j                  d	      }||||fS )
N>   r}   r|   r`   r   r   )r   r   r   r   Fr   )r
   r   r   r   r   r   )r7   r:   rj   rn   rk   rm   rl   r   r   r   rf   r   r   r   r   resultgrad_input_shapegrad_linear_weight_shapegrad_linear_bias_shaper$   r%   r&   s                         r   r#   r#     s    $ O#Ru{{5<<H!$U9,"WXX&8u{{f9v  %=CR $  kk||	J  kk##	 {{kk||	 :13CCCr   c                     | j                   }t        ||      D cg c]5  \  }}|,t        |t        j                        r|j                  |d      n|7 }}}t        |      D cg c]H  }t        t        ||      D cg c](  \  }}|t        |t        j                        r||   n|* c}} J }}}}t        j                  |D 	cg c]  }	|	d   	 c}	      }
t        j                  |D 	cg c]  }	|	d   	 c}	      }t        j                  |D 	cg c]  }	|	d   	 c}	      }t        j                  |D 	cg c]  }	|	d   	 c}	      }|
|||fdfS c c}}w c c}}w c c}}}w c c}	w c c}	w c c}	w c c}	w )zvmap rule (slow path: per-sample Python loop). A fold-into-num_batches
    fast path is now possible for both ops via the reduction='none' chunked op
    (``_linear_cross_entropy_batch_chunked_no_reduction``); left as a follow-up.
    r   r         )r   r   r   r   )	
batch_sizezip
isinstancer
   r^   movedimr   r   stack)infoin_dimsargsr	  argin_dim
moved_argsioutputsolossesgrad_inputsgrad_linear_weightsgrad_linear_biasess                 r   _vmapr    s{    J tW-	 C *S%,,"? 	FA	J  z"   	, $'z7#;C !,C1NATWW	
G  [[01!A$01F++W5qt56K++W&=qt&=>G%<qad%<=K!46HI<WW) 15&=%<s5   :E(E%-E.E%E,2E1E6>E;E%c                    d gt         z  }| j                  rT| j                  r| j                  |z  |d<   n5| j                  d c}| _        |t	        d      |j                  |      |d<   | j                  rT| j                  r| j                  |z  |d<   n5| j                  d c}| _        |t	        d      |j                  |      |d<   | j                  r^| j                  r| j                  |z  |d<   t        |      S | j                  d c}| _	        |t	        d      |j                  |      |d<   t        |      S )Nr   zjlinear_cross_entropy chunked backward called twice; retain_graph=True / double backward are not supported.r   r  )_NUM_OP_INPUTSr   r   r   r   r   r   r   r   r   r   )	r    grad_output_gi_grad_gw_grad_gb_gradr  gigwgbs	            r   ,_linear_cross_entropy_batch_chunked_backwardr%  4  sM    Vn$F
!!+-F1I''4KBz"M  ,F1I
%%!!+-F1I''4KBz"M  ,F1I ##!!+-F1I = ''4KBz"M  ,F1I=r   )setup_contextz:torch_nn::_linear_cross_entropy_batch_chunked_no_reductionc	                 :    t        | ||||d|d|||ddd      d   S )aC  reduction='none' chunked forward: returns per-sample loss ``(N,)``.

    Gradients are not precomputed (the per-sample upstream grad is
    unknown at forward); the accumulator's loss-only branch runs here,
    and the backward recomputes the chunked grads. See
    ``_linear_cross_entropy_batch_chunked_accumulator``.
    r   r   F)r   r   r   r   )r   	r7   r:   rj   rn   rk   rl   r   r   rf   s	            r   0_linear_cross_entropy_batch_chunked_no_reductionr)  p  sE    . ; #(!& 	 	r   c           	          |\	  }}}}}}}	}
}| j                  |j                         |j                         |||j                         nd ||j                         nd        || _        |	| _        |
| _        || _        y r   )save_for_backwardr   _ignore_index_batch_chunk_size_acc_policy
_acc_dtype)r    r!   r"   r7   r:   rj   rn   rk   rl   r   r   rf   s               r   >_linear_cross_entropy_batch_chunked_no_reduction_setup_contextr0    s     	
  + 7T!-4 %C,C COCNr   c	                 v    t        j                  | j                  d d | j                  | j                  d      S )Nr   Fr   )r
   r   r   r   r   r(  s	            r   r#   r#     s0     ;;BQu{{5<<u r   c                    | j                   }t        ||      D cg c]5  \  }}|,t        |t        j                        r|j                  |d      n|7 }}}t        |      D cg c]H  }t        t        ||      D cg c](  \  }}|t        |t        j                        r||   n|* c}} J }}}}t        j                  |      dfS c c}}w c c}}w c c}}}w )zvmap rule (slow path: per-sample Python loop, mirrors the
    scalar-reduction op). Stacks the per-sample ``(N,)`` losses.
    r   )	r	  r
  r  r
   r^   r  r   r)  r  )	r  r  r  r	  r  r  r  r  r  s	            r   _vmap_no_reductionr3    s    
 J
 tW-	 C *S%,,"? 	FA	J  z"   	9 $'z7#;C !,C1NATWW	
G  ;;w""s   :C(C-C.CCc                 d   | j                   \  }}}}}| j                  }|d   }|d   }	|d uxr |d   }
d gdz  }|s|	s|
st        |      S t        |||||d| j                  d| j
                  | j                  | j                  ||	|
|      \  }}}}|r||d<   |	r||d<   |
r||d<   t        |      S )Nr   r   r  	   r   r   r   )saved_tensorsneeds_input_gradr   r   r,  r-  r.  r/  )r    r  r7   r:   rj   rn   rk   needsr   r   r   r  r#   r$   r%   r&   s                   r   9_linear_cross_entropy_batch_chunked_no_reduction_backwardr9    s     9<8I8I5E=&+v  Eq!&q*$6C58)-
F8<TV} 	8!!OONN&$(	
 8Az%'7& q	!&q	$q	=r   )Tr   )!dataclasses	functoolsr   r
   __all__r   r   r'   	dataclassr)   r*   r^   r   r]   r   r   r   r   r   library	custom_opr   r  register_faker#   register_vmapr  r%  register_autogradr)  r0  r3  r9  r`   r   r   <module>rC     s    %  N4 Z Z Zz v v vP -1p<<p<<p LLp $	p
 LL4p p p p p p {{p p !%p #p llT)p  5<<u||U\\AB!p| 3"  A<<A<<A LLA $	A
 LL4A A A A A A {{A A A !%A #A  5<<u||U\\AB!AAN  %227D 7D 37Dt %22X 3X<(V $ 5 50C 6 " @r  #	<<#	<<#	 LL#	 $	#	
 LL4#	 #	 #	 #	 {{#	 \\#	#	L@ 2??  @  2??# @#.+\ 1 B B=P C r   