
    %ɷi                     *   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	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 dd
lmZmZ  e       rd dlmZ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% d dl&m'Z'm(Z( d dl)m*Z* d dl+m,Z, d dl-m.Z. d dl/m0Z0m1Z1 d dl2m3Z3 d dl4m5Z5m6Z6m7Z7m8Z8 d dl9m:Z:m;Z;m<Z<m=Z=m>Z> d dl?m@Z@mAZAmBZB d dlCmDZD d dlEmFZF d dlGmHZHmIZImJZJmKZKmLZLmMZM d dlNmOZO d dlPmQZQmRZRmSZSmTZTmUZUmVZV d dlWmXZXmYZYmZZZ d<d Z[d! Z\ G d" d#      Z]d$ Z^d% Z_ G d& d'e      Z`d( Za G d) d*      Zb G d+ d,e      Zcd- Zd G d. d/e      Ze G d0 d1ee      Zf G d2 d3ee      Zg G d4 d5ee      Zhd6 Zid=d7Zj G d8 d9ej                  j                        Zld: Zmd; Zny)>    N)ABC)partial)BCEWithLogitsLossCrossEntropyLossMSELoss   )AcceleratedOptimizer)AcceleratedScheduler   )is_megatron_lm_available)recursively_applysend_to_device)mputensor_parallel)DistributedDataParallel)finalize_model_grads)	ModelType)get_num_microbatches)get_megatron_optimizer)get_tensor_model_parallel_group"get_tensor_model_parallel_src_rank)get_forward_backward_func)get_model_config)build_train_valid_test_datasets)	BertModelT5Model)Classification)get_argsget_tensorboard_writerget_tokenizerprint_rank_last)_add_data_args_add_validation_args!core_transformer_config_from_args
parse_argsvalidate_args)load_args_from_checkpointload_checkpointsave_checkpoint)set_global_variables)gpt_builder)_compile_dependencies_init_autoresume_initialize_distributed_set_random_seedset_jit_fusion_optionswrite_args_to_tensorboard)_vocab_size_with_padding)%build_train_valid_test_data_iteratorsget_optimizer_param_schedulernum_floating_point_operationssetup_model_and_optimizer
train_steptraining_log))average_losses_across_data_parallel_groupcalc_params_l2_normget_ltor_masks_and_position_idsc           	      F   t               }|j                  rdnd}|j                  dk(  r't        d|j                   d| d       t        d       t        |      }|j                  dk(  rU|j                  r-|j                  rd	nd}t        |||j                  d
| |      }|S t        ||j                  d	| |      }|S |j                  dk(  rd|_
        t        || |dd      }|S |j                  dk(  rt        |dd
| |||      }|S t        d|j                         )zBuild the model.zpre-trainingzfine-tuningr   z	Building z model in the z mode.zThe Megatron LM model weights are initialized at random in `accelerator.prepare`. Please use `accelerator.load_checkpoint` to load a pre-trained checkpoint matching the distributed setup.bertr   T)confignum_tokentypesadd_binary_headparallel_outputpre_processpost_process)r>   num_classesr?   rB   rC   gptFN)vp_stager>   t5)r>   r?   rA   rB   rC   add_encoderadd_decoderUnsupported model type: )r   pretraining_flagrankprintmodel_type_namer$   bert_binary_headr   r   
num_labelsuse_legacy_modelsr+   r   
ValueError)	rB   rC   rH   rI   argsmoder>   r?   models	            O/var/www/html/venv/lib/python3.12/site-packages/accelerate/utils/megatron_lm.pymodel_provider_funcrW   U   sU   :D!22>DyyA~	$../~dV6JKx	
 /t4Fv%  "&"7"7QQN- $ 5 5 $')E@ L/ # OO ')E. L! 
			&!&D+|dSWX L 
			% #%##
 L 3D4H4H3IJKK    c                    | j                  d       t               }| j                  j                  j                  | j                  j                  j
                  t        d      | j                  j                  j
                  }| j                  j                  j	                  |      }t        | |      }t        | |d       }nt        j                  }|j                  dk(  rt        j                  }t        }| j                  j                  j
                   | j                  j                  j
                  }t        ||      \  }}}t        |      |_        |||fS )Nz#Preparing model optimizer schedulerzaYou must provide a `custom_model_provider_function` when using a `custom_prepare_model_function`.)	schedulerrG   )rM   r   statemegatron_lm_plugincustom_prepare_model_functioncustom_model_provider_functionrR   prepare_optimizerprepare_schedulerr   encoder_or_decoderrN   encoder_and_decoderrW   r6   len	model_len)acceleratorrS   custom_model_provider_funcrU   	optimizerrZ   
model_typemodel_provider_func_s           rV   !prepare_model_optimizer_schedulerrj      s4   ;<:D++IIU//NNVs  &1%6%6%I%I%h%h"!!44RRSmn%k59	%k9M	11
4'"66J2//NNZ#.#4#4#G#G#f#f (A )
%	9 ZDN)Y&&rX   c                   (    e Zd ZdZd Zd Zd Zd Zy)MegatronLMDummyDataLoaderz
    Dummy dataloader presents model parameters or param groups, this is primarily used to follow conventional training

    Args:
        **dataset_kwargs: Megatron data arguments.
    c                     t        j                         }t        |      }t        |      }|j	                         }t        |d         | _        | j                  j                  |       d| j                  d<   y )Nr   Tmegatron_dataset_flag)argparseArgumentParserr"   r#   parse_known_argsvarsdataset_argsupdate)selfdataset_kwargsparser	data_argss       rV   __init__z"MegatronLMDummyDataLoader.__init__   sh    ((*'%f-++-	 1.  05912rX   c                     t               }| j                  j                         D ];  \  }}t        ||d      }||k7  rt	        d| d| d| d|        t        |||       = y )N z<WARNING: MegatronLMDummyDataLoader overriding arguments for : with )r   rs   itemsgetattrrM   setattr)ru   rS   keyvalue	old_values        rV   set_megatron_data_argsz0MegatronLMDummyDataLoader.set_megatron_data_args   s~    z++113 	&JCc2.IE!RSVRWWXYbXccijminnopuovw D#u%	&rX   c                 x   d }|j                   j                  j                   |j                   j                  j                  S 	 t               }|j                  dk(  rddlm} d|_        |S |j                  dk(  rddlm} d|_        |S |j                  dk(  rddl	m} d|_        |S 	 |S # t        $ r Y |S w xY w)Nc                 N   t               }t        |j                  t        t        f      r|j                  n|j                  g|j
                  | |j                  d}|j                  dk(  r)|j                  |j                  |j                  d       n~|j                  dk(  r|j                  d|j                  i       nQ|j                  dk(  r*|j                  |j                  |j                  dd       nt        d|j                         t        d	i |\  }}}|||fS )
z&Build train, valid, and test datasets.)data_prefixsplits_stringtrain_valid_test_num_samplesseedr=   )max_seq_lengthbinary_headrE   r   rG   )r   max_seq_length_decdataset_typerJ    )r   
isinstance	data_pathlisttuplesplitr   rN   rt   
seq_lengthrO   encoder_seq_lengthdecoder_seq_lengthrR   r   )train_val_test_num_samplesrS   rs   train_dsvalid_dstest_dss         rV   "train_valid_test_datasets_providerzlMegatronLMDummyDataLoader.get_train_valid_test_datasets_provider.<locals>.train_valid_test_datasets_provider   s   :D1;DNNTSXM1Zt~~aeaoao`p!%0J			L ##v-##*.//'+'<'< %%.##($//
 %%-##*.*A*A.2.E.E(, !#;D<P<P;Q!RSS*I*YL*Y'HhXw..rX   r=   r   )r   TrE   rG   )r[   r\   *custom_megatron_datasets_provider_functionr   rN   pretrain_bertr   is_distributedpretrain_gptpretrain_t5ImportError)ru   re   r   rS   s       rV   &get_train_valid_test_datasets_providerz@MegatronLMDummyDataLoader.get_train_valid_test_datasets_provider   s    !	/F //ZZf$$77bbb	:D##v-LDH2A99%%.KDH2A99%%-JDH2A99	 . 21  	11	s   'B, -B, B, ,	B98B9c                 t   t               }| j                  |      }|j                  ~g }g }g }t        t	        |dd            D ]^  }t        j                  |       t        |      }|j                  |d          |j                  |d          |j                  |d          ` nt        |      \  }}}|||fS )Nrd   r   r   r   )	r   r   $virtual_pipeline_model_parallel_sizeranger   r   (set_virtual_pipeline_model_parallel_rankr3   append)	ru   re   rS   !train_valid_test_dataset_providertrain_data_iteratorvalid_data_iteratortest_data_iteratori	iteratorss	            rV   r3   z?MegatronLMDummyDataLoader.build_train_valid_test_data_iterators   s    z,0,W,WXc,d)44@"$"$!#74a89 8<<Q?ABcd	#**9Q<8#**9Q<8")))A,78 Lq1LH!46H #$79KKKrX   N)__name__
__module____qualname____doc__ry   r   r   r3   r   rX   rV   rl   rl      s    :&:2xLrX   rl   c                      G d d      }|d u }t        j                  |t         j                  | j                        }t         j                  j                  |t               t                      |s	|r |       S |S )Nc                       e Zd Zd Zd Zy)?_handle_megatron_data_iterator.<locals>.DummyMegatronDataloaderc                     | S Nr   ru   s    rV   __iter__zH_handle_megatron_data_iterator.<locals>.DummyMegatronDataloader.__iter__  s    KrX   c                     i S r   r   r   s    rV   __next__zH_handle_megatron_data_iterator.<locals>.DummyMegatronDataloader.__next__  s    IrX   N)r   r   r   r   r   r   rX   rV   DummyMegatronDataloaderr     s    		rX   r   dtypedevicegroup)torchtensorboolr   distributed	broadcastr   r   )re   data_iteratorr   is_data_iterator_emptyis_src_data_iterator_emptys        rV   _handle_megatron_data_iteratorr     sw      +d2!&.DEJJ_j_q_q!r	"$F$HPoPq    &*@&((rX   c           
      8   | j                  d       t               }|j                  s0ddlm}m} |j                  |j                  z  }|D ci c]  }|t        ||||          }}|d   Pt        |d   t        j                  j                  j                        r||d   _        n|d= |d= |d= ||d   _        n|d= ||d<   t        j                  j                  j                  |j                   fi |} ||| j"                  t%        j&                         t%        j(                         dd	| j*                  j-                         | j.                  
      S |j0                   |j0                  \  |_        |_        |_        nd\  |_        |_        |_        |j                  |j                  z  |_        |j9                  |       \  }}	}
|j                  |j                  z  |_        t;        | |      }t;        | |	      }	t;        | |
      }
||	|
fS c c}w )NzPreparing dataloaderr   )_PYTORCH_DATALOADER_KWARGSprepare_data_loader
batch_sizesamplershufflebatch_samplerFT)num_processesprocess_indexsplit_batchesput_on_device	rng_typesdispatch_batches)r   r   r   )re   r   )rM   r   rn   data_loaderr   r   micro_batch_sizenum_micro_batchesr   r   r   utilsdataBatchSamplerr   
DataLoaderdatasetr   r   get_data_parallel_world_sizeget_data_parallel_rankr   copyr   consumed_samplesconsumed_train_samplesconsumed_valid_samplesconsumed_test_samplesr3   r   )re   
dataloaderrS   r   r   r   kkwargsr   r   r   s              rV   r   r   !  s,   ,-:D%%Q0043I3IITnoq!WZ,Fq,IJJoo,'&+U[[-=-=-J-JK/?y!,9%9%<(5E'2'#3F< [[%%001C1CNvN
 #::<446!++002(99	
 		
   ,
 %%	++* dk`D')DdF` $ 5 58N8N N <<[I		
 $ 5 59O9O O<#3F
 =#3F
 <cuv"$79KKKo ps   Hc                   <     e Zd Z fdZddZd Zed        Z xZS )MegatronLMOptimizerWrapperc                 *    t         |   |dd        y )NF)device_placementscalersuperry   )ru   rg   	__class__s     rV   ry   z#MegatronLMOptimizerWrapper.__init__d  s    U4HrX   c                      y r   r   )ru   set_to_nones     rV   	zero_gradz$MegatronLMOptimizerWrapper.zero_gradg      rX   c                      y r   r   r   s    rV   stepzMegatronLMOptimizerWrapper.stepj  r   rX   c                 .    | j                   j                  S )zTWhether or not the optimizer step was done, or skipped because of gradient overflow.)rg   skipped_iterr   s    rV   step_was_skippedz+MegatronLMOptimizerWrapper.step_was_skippedm  s     ~~***rX   r   )	r   r   r   ry   r   r   propertyr   __classcell__r   s   @rV   r   r   c  s'    I + +rX   r   c                     | j                  d       t               }t        ||j                  |j                  |j
                        S )NzPreparing optimizer)rM   r   r   no_wd_decay_condscale_lr_condlr_mult)re   rU   rS   s      rV   r_   r_   s  s<    +,:D!%)>)>@R@RTXT`T`aarX   c                       e Zd ZdZddZy)MegatronLMDummySchedulera  
    Dummy scheduler presents model parameters or param groups, this is primarily used to follow conventional training
    loop when scheduler config is specified in the deepspeed config file.

    Args:
        optimizer (`torch.optim.optimizer.Optimizer`):
            The optimizer to wrap.
        total_num_steps (int):
            Total number of steps.
        warmup_num_steps (int):
            Number of steps for warmup.
        **kwargs (additional keyword arguments, *optional*):
            Other arguments.
    Nc                 <    || _         || _        || _        || _        y r   )rg   total_num_stepswarmup_num_stepsr   )ru   rg   r  r  r   s        rV   ry   z!MegatronLMDummyScheduler.__init__  s     ". 0rX   Nr   )r   r   r   r   ry   r   rX   rV   r  r  z  s    rX   r  c                   $     e Zd Z fdZd Z xZS )MegatronLMSchedulerWrapperc                 &    t         |   ||       y r   r   )ru   rZ   
optimizersr   s      rV   ry   z#MegatronLMSchedulerWrapper.__init__  s    J/rX   c                      y r   r   )ru   rS   r   s      rV   r   zMegatronLMSchedulerWrapper.step  s    rX   )r   r   r   ry   r   r   r   s   @rV   r	  r	    s    0rX   r	  c                 >    | j                  d       t        |      }|S )NzPreparing scheduler)rM   r4   )re   rg   rZ   s      rV   r`   r`     s!    +,-i8IrX   c                   4     e Zd ZdZ fdZd Zd Zd Z xZS )AbstractTrainStepz;Abstract class for batching, forward pass and loss handler.c                 0    t         |           || _        y r   )r   ry   name)ru   r  r   s     rV   ry   zAbstractTrainStep.__init__  s    	rX   c                      y r   r   )ru   re   rn   s      rV   get_batch_funcz AbstractTrainStep.get_batch_func  r   rX   c                      y r   r   r   s    rV   get_forward_step_funcz'AbstractTrainStep.get_forward_step_func  r   rX   c                      y r   r   )ru   re   s     rV   get_loss_funczAbstractTrainStep.get_loss_func  r   rX   )	r   r   r   r   ry   r  r  r  r   r   s   @rV   r  r    s    ErX   r  c                   4     e Zd ZdZ fdZd Zd Zd Z xZS )BertTrainStepzg
    Bert train step class.

    Args:
        args (`argparse.Namespace`): Megatron-LM arguments.
    c                 V   t         |   d       | j                  ||j                        | _        | j                  ||j                  |j                        | _        | j                  |j                  |j                        | _        |j                  sd | _        y ddlm} || _        y )Nr  r   )SequenceClassifierOutput)r   ry   r  rn   	get_batchr  rK   rP   	loss_funcr  rO   forward_stepmodel_return_dictmodel_output_classtransformers.modeling_outputsr  )ru   re   rS   r  r   s       rV   ry   zBertTrainStep.__init__  s    ),,[$:T:TU++K9N9NPTP_P_` 66t7L7LdNcNcd%%&*D#N&>D#rX   c                     d }d }|j                   j                  j                   |j                   j                  j                  S |r		 ddlm} |S |S # t
        $ r Y |S w xY w)Nc                 l   g d}t         j                  }| t        |       }nd}t        j                  |||      }|d   j                         }|d   j                         }|d   j                         }|d   j                         }|d   j                         }	|d   j                         }
|||||	|
fS )	Build the batch.)texttypeslabels	is_random	loss_maskpadding_maskNr%  r&  r(  r)  r'  r*  r   int64nextr   broadcast_datalongfloat)r   keysdatatyper   data_btokensr&  sentence_orderr)  	lm_labelsr*  s              rV   get_batch_megatronz8BertTrainStep.get_batch_func.<locals>.get_batch_megatron  s     YD{{H (M*$33D$IF F^((*F7O((*E#K0557N{+113Ix(--/I!.1668L5.)YTTrX   c                    t        |       }t        |t        j                  j	                               }|d   j                         }|d   j                         }d|v r|d   j                         }nd}d|v r9|d   j                         }|d   dk7  j                  t        j                        }nd}d}d|v r|d   j                         }nd}||||||fS )r$  	input_idsattention_masktoken_type_idsNr'  next_sentence_label)r-  r   r   cudacurrent_devicer/  tor0  )r   r   r4  r*  r&  r6  r)  r5  s           rV   get_batch_transformerz;BertTrainStep.get_batch_func.<locals>.get_batch_transformer  s    &D!$

(A(A(CDD +&++-F 01668L4'-.3354 N//1	!(^t377D	 	 	$,!%&;!<!A!A!C!%5.)YTTrX   r   r  )r[   r\   custom_get_batch_functionr   r  r   ru   re   rn   r7  rA  r  s         rV   r  zBertTrainStep.get_batch_func  ss    	U0	U2 //IIU$$77QQQ 3  
 )(	  %%   
A 	A! A!c                      d } fd}|j                   j                  j                   |j                   j                  j                  S |r|S |S )Nc                    |\  }}|j                         }| j                         } t        j                  |j                  d      | j	                  d      z        | j                         z  }|tt        j                  |j                  dd      j                         |j                  d      d      }|j                         }||z   }t        ||g      }||d   |d   dfS |}t        |g      }|d|d   ifS )Nr   )ignore_indexr   r   )lm losszsop lossrJ  )r0  r   sumviewreshapeFcross_entropyr9   )	r)  r5  output_tensorlm_loss_
sop_logitslm_losssop_losslossaveraged_lossess	            rV   loss_func_pretrainz7BertTrainStep.get_loss_func.<locals>.loss_func_pretrain  s    #0 Hj~~'H!)Iiib 1I4E4Eb4I IJY]]_\G%??:??2q+A+G+G+I>K^K^_aKbqst#>>+)"KWV^L_"`);YZI[\\\ "KWI"Vi);<<<rX   c                    dk(  r2t               } ||j                  d      | j                  d            }nj                  dkD  r_| j                  t        j
                  t        j                  fv r3t               } ||j                  d      | j                  d            }nt               } |||       }t        |g      }|d|d   ifS )Nr   rH  rU  r   )
r   rL  rP   r   r   r/  intr   r   r9   )r'  logitsloss_fctrU  rV  rP   ru   s        rV   loss_func_finetunez7BertTrainStep.get_loss_func.<locals>.loss_func_finetune  s    Q"9BRA1$&,,5::uyy:Q*Q+-B
 ;V[[_M,./GOO&/!"4555rX   r[   r\   custom_loss_function)ru   re   rK   rP   rW  r\  s   `  `  rV   r  zBertTrainStep.get_loss_func  sN    	=&	6 //DDP$$77LLL%%%%rX   c                       fd}|S )Nc                     j                  |       \  }}}}}}
sd}r% |||||      }|t        j                  ||      fS  ||||      }	|	t        j                  |      fS )Forward step.Ntokentype_idsr6  )rc  r  r   r  )r   rU   r4  r&  r5  r)  r'  r*  rP  rZ  rO   rK   ru   s             rV   r  z9BertTrainStep.get_forward_step_func.<locals>.forward_step.  sw    MQ^^\iMjJFE>9fl# %fl%[a b$gdnni&XXXv|5Iwt~~v>>>rX   r   )ru   rK   rO   r  s   ``` rV   r  z#BertTrainStep.get_forward_step_func-  s    	? rX   	r   r   r   r   ry   r  r  r  r   r   s   @rV   r  r    s    
?>)@'&RrX   r  c                   4     e Zd ZdZ fdZd Zd Zd Z xZS )GPTTrainStepzf
    GPT train step class.

    Args:
        args (`argparse.Namespace`): Megatron-LM arguments.
    c                    t         |   d       | j                  ||j                        | _        | j                  |      | _        | j                         | _        |j                  t               }|j                  | _        |j                  | _        |j                  | _        |j                  | _        |j                   | _        |j"                  | _        |j$                  sd | _        y ddlm} || _        y )Nrg  r   )!CausalLMOutputWithCrossAttentions)r   ry   r  rn   r  r  r  r  r  
vocab_filer    eod	eod_tokeneos_token_id	pad_tokenreset_position_idsreset_attention_maskeod_mask_lossr  r   r!  ri  )ru   re   rS   	tokenizerri  r   s        rV   ry   zGPTTrainStep.__init__F  s    (,,[$:T:TU++K8 668??&%I&]]DN****"&"9"9$($=$=!!//%%&*D#W&GD#rX   c                       fd} fd}|j                   j                  j                   |j                   j                  j                  S |r		 ddlm} |S |S # t
        $ r Y |S w xY w)Nc           	         dg}t         j                  }| t        |       }nd}t        j                  |||      }|d   j                         }|ddddf   j                         }|ddddf   j                         }t        |j                  j                  j                  j                  j                  d      \  }}	}
|||	||
fS )zGenerate a batchr%  Nr   rH  Trl  rn  ro  rp  rq  pad_mask_loss)r   r,  r-  r   r.  r/  
contiguousr;   rl  ro  rp  rq  )r   r1  r2  r   r3  tokens_r'  r4  r:  r)  position_idsru   s              rV   r7  z7GPTTrainStep.get_batch_func.<locals>.get_batch_megatron[  s     8D{{H (M*$33D$IF Vn))+GQU^..0FQV_//1F 7V....#'#:#:%)%>%>"00"73NI| 69nlJJrX   c           	      b   t        |       }d|d   i}t        |t        j                  j	                               }|d   j                         }t        j                  |j                  d   df|j                  |j                        	j                  z   }t        j                  ||gd      }|d d dd f   j                         }|d d d df   j                         }t        |	j                  	j                  	j                  	j                  	j                   d      \  }}}|||||fS )	Nr9  r   r   r   dimrH  Tru  )r-  r   r   r>  r?  r/  zerosshaper   r   rl  concatrw  r;   ro  rp  rq  )
r   r   rx  paddingr'  r4  r:  r)  ry  ru   s
            rV   rA  z:GPTTrainStep.get_batch_func.<locals>.get_batch_transformery  s   &Dk!23D!$

(A(A(CDD;',,.Gkk7==#3Q"7w}}U\UcUcdgkguguuGllGW#51=GQU^..0FQV_//1F6U....#'#:#:%)%>%>"00"73NI| 69nlJJrX   r   rB  )r[   r\   rC  r   r  r   rD  s   `     rV   r  zGPTTrainStep.get_batch_funcZ  st    	K<	K, //IIU$$77QQQ 2  
 )(	  %%s   A 	A&%A&c                     t               fd}|j                  j                  j                   |j                  j                  j                  S |S )Nc                    j                   r|\  }}n|}|j                         }| j                  d      j                         } j                  dkD  rt	        j
                  t	        j                  |j                  d      | z        j                  d      | j                         j                  d      g      }t        j                  j                  |t        j                                |d   |d   z  }n8t	        j                  |j                  d      | z        | j                         z  }j                  rot        j                  j                         }|j                         rAJ d| dt        j                  j                          dt!        j"                         d           t%        |g      }d|d   i}j                   r|j'                  d	i       ||fS )
NrH  r   r   r   zRank z7: found NaN in local forward loss calculation. Device: z, node: rJ  rZ  )return_logitsr0  rL  context_parallel_sizer   catrK  r   
all_reducer   get_context_parallel_groupcheck_for_nan_in_loss_and_gradget_rankisnanr>  r?  osunamer9   rt   )	r)  rP  lossesrZ  rU  global_rankaveraged_lossoutput_dictrS   s	           rV   r  z-GPTTrainStep.get_loss_func.<locals>.loss_func  s   !!!.&\\^F!r*002I))A-yy%))FKKOi,G"H"M"Ma"PR[R_R_RaRfRfghRi!jk!!,,T9W9W9Y,ZAwa(yyR9!<=	O 22#//88:::< K= )$zz88:;8BHHJqM?T' FtfMM$mA&67K!!""Hf#56$$rX   )r   r[   r\   r^  )ru   re   r  rS   s      @rV   r  zGPTTrainStep.get_loss_func  sG    z	%< //DDP$$77LLLrX   c                       fd}|S )Nc                 z    j                  |       \  }}}}} |||||      }|t        j                  |      fS )ra  )r'  rd  )	r   rU   r4  r'  r)  r:  ry  rP  ru   s	           rV   r  z8GPTTrainStep.get_forward_step_func.<locals>.forward_step  sG     GKnnUbFcCFFI~|!&,vVM '$..)"DDDrX   r   ru   r  s   ` rV   r  z"GPTTrainStep.get_forward_step_func  s    	E rX   re  r   s   @rV   rg  rg  >  s     H(A)F#J	rX   rg  c                   d     e Zd ZdZ fdZed        Zed        Zed        Zd Z	d Z
d Z xZS )	T5TrainStepze
    T5 train step class.

    Args:
        args (`argparse.Namespace`): Megatron-LM arguments.
    c                     t         |   d       | j                  ||j                        | _        | j                  |      | _        | j                         | _        |j                  sd | _
        y ddlm} || _
        y )Nr  r   )Seq2SeqLMOutput)r   ry   r  rn   r  r  r  r  r  r  r   r!  r  )ru   re   rS   r  r   s       rV   ry   zT5TrainStep.__init__  si    ',,[$:T:TU++K8 668%%&*D#E&5D#rX   c                 ^    | j                  d      }| j                  d      }||z  }|dk  }|S )Nr   r         ?)	unsqueeze)r:  attention_mask_b1sattention_mask_bs1attention_mask_bssextended_attention_masks        rV   attn_mask_postprocessz!T5TrainStep.attn_mask_postprocess  sC     ,55a8+55a8/2DD"4s":&&rX   c                 j    t        j                  t        j                  d| | f|            }|dk  }|S Nr   r   r  )r   trilones)r   r   r:  s      rV   get_decoder_maskzT5TrainStep.get_decoder_mask  s3    EJJ:z/JSY$Z['#-rX   c                     | j                   \  }}| j                  d      }t        j                  ||df|      }||z  }|dk  }|S r  )r~  r  r   r  )	r:  dec_seq_lengthr   r   _r  r  r  r  s	            rV   get_enc_dec_maskzT5TrainStep.get_enc_dec_mask  sZ    &,,
A ,55a8"ZZ^Q(GPVW/2DD"4s":&&rX   c                     d }d }|j                   j                  j                   |j                   j                  j                  S |r		 ddlm} |S |S # t
        $ r Y |S w xY w)Nc                 R   g d}t         j                  }| t        |       }nd}t        j                  |||      }|d   j                         }|d   j                         }|d   j                         }|d   j                         }|d   dk  }	|d	   dk  }
|d
   dk  }|||||	|
|fS )r$  )text_enctext_decr'  r)  enc_maskdec_maskenc_dec_maskNr  r  r'  r)  r  r  r  r  r+  )r   r1  r2  r   r3  
tokens_enc
tokens_decr'  r)  r  r  r  s               rV   r7  z6T5TrainStep.get_batch_func.<locals>.get_batch_megatron  s     kD{{H (M*$33D$IF  
+002J
+002JH%**,F{+113Ij)C/Hj)C/H!.1C7Lz9fhR^^^rX   c                 :   t        |       }t        |t        j                  j	                               }|d   j                         }|d   j                         }|dk7  j                  t        j                        }d|v r|d   j                         }nn|j                  |j                  |j                  t        j
                        }|dddf   j                         |dd	df<   d
|d<   |j                  |dk(  d
       t        j                  |d   j                               }t        j                  |j                  d	   |j                        }t        j!                  |d   j                         |j                  d	   |j                        }|||||||fS )r$  r9  r'  r<  decoder_input_ids)r   r   .NrH  r   r   ).r   r:  )r-  r   r   r>  r?  r/  r@  r0  	new_zerosr~  r   clonemasked_fill_r  r  r  r  )	r   r   r  r'  r)  r  r  r  r  s	            rV   rA  z9T5TrainStep.get_batch_func.<locals>.get_batch_transformer  sz   &D!$

(A(A(CDDk*//1J(^((*F4++EKK8I"d*!"56;;=
#--fll6==X]XbXb-c
&,S#2#X&6&<&<&>
37#%&
6"''
d(:A>"88>N9O9T9T9VWH"33J4D4DQ4GIZIZ[H&77%&++-z/?/?/BJDUDUL z9fhR^^^rX   r   rB  )r[   r\   rC  r   r  r   rD  s         rV   r  zT5TrainStep.get_batch_func  ss    	_2	_. //IIU$$77QQQ 1  
 )(	  %%rE  c                     d }|j                   j                  j                   |j                   j                  j                  S |S )Nc                     |j                         }t        j                  |j                  d      | j	                  d      z        | j                         z  }|}t        |g      }|d|d   ifS )NrH  rJ  r   )r0  r   rK  rL  rM  r9   )r)  rP  rQ  rS  rU  rV  s         rV   r  z,T5TrainStep.get_loss_func.<locals>.loss_funcA  sh    $**,Hiib 1I4E4Eb4I IJY]]_\GDG	RO)_Q%7888rX   r]  )ru   re   r  s      rV   r  zT5TrainStep.get_loss_func@  s?    	9 //DDP$$77LLLrX   c                       fd}|S )Nc           	          
j                  |       \  }}}}}}} ||||||d|      }	|	t        
j                  |      fS )ra  Nrb  rd  )r   rU   r  r  r)  r6  r  r  r  rP  ru   s             rV   r  z7T5TrainStep.get_forward_step_func.<locals>.forward_stepO  s_     ^b]k]k^ZJ
Iy(Hl "J(LX\hqM !'$..)"DDDrX   r   r  s   ` rV   r  z!T5TrainStep.get_forward_step_funcN  s    	E rX   )r   r   r   r   ry   staticmethodr  r  r  r  r  r  r   r   s   @rV   r  r    sY    
6 
' 
'  
 	' 	'=)~rX   r  c                      t               } t        d d d        | j                  dk(  rt        d| j                   d       t        | j                  | j                         y )Nr   z> setting random seeds to z ...)r   r.   rL   rM   r   r/   data_parallel_random_init)rS   s    rV   finish_mpu_initr  _  sL    :DD$- yyA~*499+T:;TYY > >?rX   c                 D   |i }| j                  d       t        j                  j                         sJ d       t	        |d      }|j                         D ]Q  \  }}t        ||d       2|j                  dk(  r#t        d| dt        ||       d| d| d	       t        |||       S |j                  s|j                  d
d      r|j                  J d       t        |       t        |       t        |d       t                t!                t#                t%                t'               }t        |dd       t)        |j*                  |      |_        |j.                  dk(  r*|j0                  r|j2                  dk(  rd|_        d|_        y d|_        d|_        y )NzInitializing Megatron-LMzMegatron requires CUDA.T)ignore_unknown_argsr   z*WARNING: overriding default arguments for r|   r}   )flushuse_checkpoint_argsFz/--use-checkpoints-args requires --load argument)build_tokenizerpadded_vocab_sizer=   r   )rM   r   r>  is_availabler%   r~   r   rL   r   r  getloadr'   r&   r*   r  r-   r,   r0   r   r2   orig_vocab_sizer  rN   rK   rP   rO   	iteration)re   extra_args_providerargs_defaultsrS   r   r   s         rV   
initializer  l  s   01::""$?&??$ )tDD $))+ "
U4d#/yyA~@QwtUXGYFZZ`ad`eefglfmn 	c5!" =#4#45JE#Ryy$W&WW$!$'$ u5    :Dt($/7!9$:N:NPT!Uv%$*?*?DOOWXDX $ DN !&DNrX   c                   ^     e Zd ZdZ fdZd Zd Zd Zd Zd Z	d Z
d	 Zd
 Zd Zd Z xZS )MegatronEnginez
    Megatron-LM model wrapper

    Args:
        accelerator (:class:`~accelerate.Accelerator`): The accelerator object to use.
        model: Megatron-LM model
        optimizer: Megatron-LM optimizer
        lr_scheduler: Megatron-LM lr scheduler
    c                    t         |           || _        |d   | _        || _        || _        t               }|j                  j                  j                  K |j                  j                  j                  |fi |j                  j                  j                  | _        n{|j                  dk(  rt        ||      | _        nZ|j                  dk(  rt        ||      | _        n9|j                  dk(  rt        ||      | _        nt!        d|j                         d| j                  _        i | _        i | _        d| _        d| _        d| _        d | _        |j0                  t3                y y )Nr   r=   rE   rG   rJ   FT)r   ry   module
base_modelrg   rZ   r   r[   r\   custom_train_step_classcustom_train_step_kwargstrain_step_handlerrN   r  rg  r  rR   r   total_loss_dicteval_total_loss_dictr  report_memory_flag$num_floating_point_operations_so_farmodule_configtensorboard_dirr1   )ru   re   rU   rg   rZ   rS   r   s         rV   ry   zMegatronEngine.__init__  sR   (""z//GGS&bk&7&7&J&J&b&b'#))<<UU'D# !!V+&3K&FD#!!U*&2;&ED#!!T)&1+t&DD#78L8L7MNOO&+#  "$&!"&451!+%' ,rX   c                     t               }t         j                  d         } j                  j                  |_        t         j                  d   t              r|j                  r|j                  J d        j                  D cg c]  }|j                   c}|_	        t         j                        dk(  r|j                  d   |_	        |j                  rU j                  D cg c]  }|j                   c}|_        t         j                        dk(  r|j                  d   |_        |j                  rn|j                   rbt#        t         j                              D cg c]   fd
 c}|_        t         j                        dk(  r|j$                  d   |_        t&        |_        |S c c}w c c}w c c}w )Nr   zWhen overlap_grad_reduce is True, config.no_sync_func must be None; a custom no_sync_func is not supported when overlapping grad-reducer   c                 <    j                   j                  |       S r   )rg   finish_param_sync)xmodel_indexru   s    rV   <lambda>z2MegatronEngine.get_module_config.<locals>.<lambda>  s    $..::;J rX   )r   r   r  rg   
scale_lossgrad_scale_funcr   LocalDDPoverlap_grad_reduceno_sync_funcno_syncrc   delay_grad_reducestart_grad_syncgrad_sync_funcoverlap_param_gatherdelay_param_gatherr   param_sync_funcr   finalize_model_grads_func)ru   rS   r>   model_chunkr  s   `   `rV   get_module_configz MegatronEngine.get_module_config  sx   z!$++a.1!%!:!:dkk!nh/D4L4L&&. V. KO++"V;;#6#6"VF4;;1$&,&9&9!&<#%%X\XcXc(d)D)D(d%t{{#q(,2,A,A!,DF)$$)@)@^cdghlhshsdt^u&OZJ&F" 4;;1$)/)?)?)B&+?( #W )e&s   
F9+F>+Gc                     | j                   D ]  }|j                           | j                  | j                         | _        | j	                          y r   )r  trainr  r  log_eval_resultsru   model_modules     rV   r  zMegatronEngine.train  sL     KK 	!L 	! %!%!7!7!9DrX   c                     | j                   D ]  }|j                           | j                  | j                         | _        y y r   )r  evalr  r  r  s     rV   r  zMegatronEngine.eval  sE     KK 	 L	  %!%!7!7!9D &rX   c                 x   t               }g }t        |      dkD  r|j                  dkD  rot        d|j                        D ]U  }|j	                  |j                         D ci c](  \  }}||||j                  z  |dz   |j                  z   * c}}       W n|g}t        | j                        dkD  r`t        |      dkD  r7t        t        | j                              D cg c]  }t        |       c}}|S d gt        | j                        z  }|S t        |      dkD  rt        |      nd }|S c c}}w c c}w )Nr   r   )	r   rc   r   r   r   r~   r   r  iter)	ru   
batch_datarS   data_chunksr   r   vr  batch_data_iterators	            rV   get_batch_data_iteratorz&MegatronEngine.get_batch_data_iterator  sC   zz?Q%%)q$"8"89 A&& )3(8(8(: $1 qT%:%:!:a!etG\G\=\]]  *lt{{a z?Q& -2#dkk2B,CDqk"D   #"	 Vc$++..   #" 8;:7J${"3PT""! Es   !-D1"D7c           
         | j                  |      }t        | j                  j                  || j                  | j
                  | j                  | j                  t                     \  }}}}}}}|dk(  | j
                  _	        ||||fS )z
        Training step for Megatron-LM

        Args:
            batch_data (:obj:`dict`): The batch data to train on.
        )forward_step_funcr   rU   rg   opt_param_schedulerr>   forward_backward_funcr   )
r  r7   r  r  r  rg   rZ   r  r   r   )ru   r  r  loss_reducedr   r  	grad_normnum_zeros_in_grads           rV   r7   zMegatronEngine.train_step  s     #:::FLV"55BB-++nn $%%";"=M
IlAq!Y8I '3a&7#\96GGGrX   c           	         t               }| j                  |      }t               } || j                  j                  || j
                  t               |j                  |j                  d      }|j                  dk\  rt        j                  j                          |xj                  t        j                         |j                  z  t               z  z  c_        t        j                   d      rni }|d   D ]b  }|D cg c]  }||   	 }	}t#        |	d   j$                        dk(  rt'        |	      t#        |	      z  ||<   Kt        j(                  |	      ||<   d |S i S c c}w )z
        Evaluation step for Megatron-LM

        Args:
            batch_data (:obj:`dict`): The batch data to evaluate on.
        T)r   r   rU   num_microbatchesr   r   forward_onlyr   )ignore_virtualr   )r   r  r   r  r  r  r   r   r   empty_unused_memory_levelr   r>  empty_cacher   r   r   is_pipeline_last_stagerc   r~  rK  r  )
ru   r  rS   r  r  
loss_dictsr  r   r  losses_reduced_for_keys
             rV   	eval_stepzMegatronEngine.eval_step#  sO    z":::F 9 ;*"55BB-++13!22

 ))Q.JJ""$##,,.1F1FFI]I__	
# %%T:L!!} M:D)EQ!C&)E&)E-a06671<(+,B(CcJ`Fa(aL%(-5K(LL%M  	 *Fs   ?E!c                    t               }| j                  d   j                  r6 | j                  d
i |\  }}}}| xj                  dz  c_        t        j                         |j                  z  t               z  }|xj                  |z  c_	        | xj                  t        ||      z  c_
        |j                  }| j                  j                         j                         }d }	|j                   rt#        | j$                        }	t'        || j(                  | j                  j*                  d   d   | j                  || j,                  |||	|
      | _        n | j.                  d
i |}|j                  |D ]  }
| j0                  j3                  |
t4        j6                  j9                  dg            ||
   z   | j0                  |
<   | j0                  j3                  |
dz   t4        j6                  j9                  dg            t4        j6                  j9                  dg      z   | j0                  |
dz   <    t5        j:                  dt4        j6                  j=                               }|D ]&  }
t?        ||
   j@                        dk(  s|||
   z  }( d }d|v r|d   }| jB                  jD                  | jB                  jE                  ||	      S |S )Nr   r   lrg        
_num_itersg      ?r  rZ  )rU  rZ  r   )#r   r  trainingr7   r  r   r   r   r   r   r  r5   r  rg   get_loss_scaleitemlog_params_normr:   rU   r8   r  param_groupsr  r  r  r  r   r>  FloatTensorr   r?  rc   r~  r  r   )ru   r  rS   	loss_dictr   r  r  r   
loss_scaleparams_normr   rU  rZ  s                rV   forwardzMegatronEngine.forwardK  s    z;;q>""DSDOODaV`DaAI|Y0ANNaN99;d>S>SSVjVllJ'':5'559VW[]g9hh5##/!^^::<AAC
"''"5djj"AK*6((NN//248NN++ %+' '44I##/$ 6C1155c5::;Q;QSVRW;XY\efi\jj --c2 EID]D]DaDal*EJJ,B,BC5,IE

..u5E6D--cL.@A	6 ||C

(A(A(CD 	'C9S>''(A-	#&	' y x(F""55A**==4PV=WWrX   c                    t               }|j                  | j                  dk(  ry t               }t               }d| j                   d}| j                  D ]  }|j                  d      r| j                  |   | j                  |dz      z  }|| d| dz  }t        j                  t        d|j                                     }|j                  r|| d| dz  }|s|j                  | d|j                         | j                         |j                  s|j                  | d	|| j                          t        |      d
z   }t        d|z         t        |       t        d|z         i | _        y )Nr   zvalidation loss at iteration z | r  z value:    z PPL: z validationz validation pplr   -)r   r  r  r   r  endswithmathexpminr  rK   
add_scalarrc   r!   )ru   rS   writerstringr   r   ppllengths           rV   r  zMegatronEngine.log_eval_results  sf   z'4>>Q+>z')00@D,, 	TC||L)--c2T5N5NsUaOa5bbEXeWC00F((3r5::<01C$$SEuC00!!SE"5uzz|T^^T((%%_&=sDNNS	T Vqf%f%$&!rX   c                 B   | j                          t               }||_        t        j                  j                          t        | j                  | j                  | j                  | j                  | j                         t        j                  j                          y )N)r  )r  r   saver   r   barrierr)   r  r  rg   rZ   r  )ru   
output_dirrS   s      rV   r)   zMegatronEngine.save_checkpoint  so    z	!!#NNKKNNNN151Z1Z	
 	!!#rX   c                    t               }||_        d|_        d|_        t        j
                  j                          t        | j                  | j                  | j                        \  }}t        j
                  j                          || _        || _        |j                  r+| j                  dk(  r| j                  j                          y y y r  )r   r  r   r   r   r   r+  r(   r  rg   rZ   r  r  fp16reload_model_params)ru   	input_dirrS   r  r  s        rV   r(   zMegatronEngine.load_checkpoint  s    z	&'#&'#!!#:I$++W[WeWegkgugu:v7	7!!#"4X1991,NN..0 -9rX   )r   r   r   r   ry   r  r  r  r  r7   r  r  r  r)   r(   r   r   s   @rV   r  r    sB    (>4 :#2H0&P=~'4$1rX   r  c                     t        |       S )z
    Average losses across data parallel group.

    Args:
        losses (List[Tensor]): List of losses to average across data parallel group.
    )r9   )r  s    rV   %avg_losses_across_data_parallel_groupr2    s     5V<<rX   c                 $    d }t        || d      S )z
    Recursively gather tensor in a nested list/tuple/dictionary of tensors from data parallel ranks.

    Args:
        tensor (nested list/tuple/dictionary of `torch.Tensor`):
            The data to gather across data parallel ranks.

    c                    | j                   dk(  r| j                         d    } t        t        j                  j                  t        j                                     D cg c]  }t        j                  |        }}t        j                  j                  || t        j                                t        j                  |d      S c c}w )Nr   r   r{  )ndimr  r   r   r   get_world_sizer   get_data_parallel_group
empty_like
all_gatherr  )r   r  output_tensorss      rV   _gpu_gather_onez;gather_across_data_parallel_groups.<locals>._gpu_gather_one  s    ;;!\\^D)F 5,,;;#B]B]B_;`a
 V$
 
 	$$^V3C^C^C`$ayyQ//
s    C	T)error_on_other_type)r   )r   r;  s     rV   "gather_across_data_parallel_groupsr=    s    0 _f$OOrX   )TTTT)NN)oro   r!  r  abcr   	functoolsr   r   torch.nn.functionalnn
functionalrN  torch.nnr   r   r   rg   r	   rZ   r
   importsr   
operationsr   r   megatron.corer   r   megatron.core.distributedr   r  r   megatron.core.enumsr   )megatron.core.num_microbatches_calculatorr   megatron.core.optimizerr   megatron.core.parallel_stater   r   megatron.core.pipeline_parallelr   megatron.core.utilsr   "megatron.legacy.data.dataset_utilsr   megatron.legacy.modelr   r   $megatron.legacy.model.classificationr   megatron.trainingr   r   r    r!   megatron.training.argumentsr"   r#   r$   r%   r&   megatron.training.checkpointingr'   r(   r)   megatron.training.global_varsr*   megatron.training.gpt_buildersr+   megatron.training.initializer,   r-   r.   r/   r0   r1   %megatron.training.tokenizer.tokenizerr2   megatron.training.trainingr3   r4   r5   r6   r7   r8   megatron.training.utilsr9   r:   r;   rW   rj   rl   r   r   r   r_   r  r	  r`   r  r  rg  r  r  r  Moduler  r2  r=  r   rX   rV   <module>r[     sc     	      A A , , - 9 2M>-N>pI4R8C   lkB:  O  .b'8jL jLZ$>LD+!5 + b .!5  "K% K\M$ M`N# Nb	@/d_1UXX__ _1F	=PrX   