Ë
    Gêñi&  ã                   ó~   — d dl Z d dlZd dlZd dlmZ ddlmZmZmZm	Z	  e	j                  e«      Zdefd„Zd„ Zd	d„Zy)
é    N)Ú
DataLoaderé   )ÚWEIGHTS_NAMEÚPushToHubMixinÚis_torch_xla_availableÚloggingÚ
dataloaderc                 óÚ   — t        «       r`dd lmc m} t	        | |j
                  «      sJ d«       ‚dd lmc m} |j                  |j                  «       d«      }|| j                  d<   | S | S )Nr   zPThe dataloader must be a `torch_xla.distributed.parallel_loader.MpDeviceLoader`.)ÚfsdpNÚinput_sharding)r   Ú%torch_xla.distributed.parallel_loaderÚdistributedÚparallel_loaderÚ
isinstanceÚMpDeviceLoaderÚtorch_xla.distributed.spmdÚspmdÚShardingSpecÚget_global_meshÚ_parallel_loader_kwargs)r	   ÚplÚxsÚsharding_specs       ú_/var/www/pod-logistic/pod-ai/venv/lib/python3.12/site-packages/transformers/integrations/tpu.pyÚtpu_spmd_dataloaderr      sl   € ÜÔß:Ð:ä˜* b×&7Ñ&7Ô8ð 	
Ø^ó	
Ð8÷ 	0Ð/àŸ™¨×(:Ñ(:Ó(<¸nÓMˆØ?Lˆ
×*Ñ*Ð+;Ñ<ØÐàÐó    c                 ó2  ‡‡‡‡‡‡— ddl mc mŠ ddlmc mŠ ddlm} 	 ddlm	Š ddlm
Š ddlm}m} ‰rddlmŠ d}d}t#        | d
d«      }|j$                  j'                  d|«      }	|j$                  d   dkD  r%t)        j*                  ||j$                  d   ¬«      }nQ|	�Ot-        «       }
|	D ])  } || |«      }|€t/        d«      ‚|
j1                  |«       Œ+ t)        j*                  ||
¬«      }|j2                  }|j$                  d   rD| j4                  j6                  r&t8        j;                  d«       d| j4                  _        ˆˆˆˆfd„}‰rˆfd„} ‰| |||¬«      } n ‰| f||dœ|¤Ž} di fˆfd„	}|‰_        | S # t         $ r t!        d	«      ‚w xY w)a.  
    Wraps a model with XLA Fully Sharded Data Parallelism (FSDP).

    Handles both FSDP v1 (`XlaFullyShardedDataParallel`) and v2 (`SpmdFullyShardedDataParallel`),
    including auto-wrap policies, gradient checkpointing, and patching `xm.optimizer_step`.

    Args:
        model (`torch.nn.Module`): The model to wrap.
        args (`TrainingArguments`): The training arguments containing FSDP configuration.
        is_fsdp_xla_v2_enabled (`bool`): Whether FSDP v2 (SPMD) is enabled.

    Returns:
        `torch.nn.Module`: The FSDP-wrapped model.
    r   Nr   )Úget_module_class_from_name)ÚXlaFullyShardedDataParallel)Úcheckpoint_module)Úsize_based_auto_wrap_policyÚtransformer_auto_wrap_policy)ÚSpmdFullyShardedDataParallelzJMissing XLA FSDP related module; please make sure to use torch-xla >= 2.0.Ú_no_split_modulesÚtransformer_layer_cls_to_wrapÚmin_num_params)r&   z@Could not find the transformer layer class to wrap in the model.)Útransformer_layer_clsÚxla_fsdp_grad_ckptzX`use_cache=True` is incompatible with gradient checkpointing. Setting `use_cache=False`.Fc                 ó4   •— ‰s‰n‰} | ‰| «      g|¢­i |¤ŽS ©N© )ÚmÚargsÚkwargsÚ
target_clsÚFSDPÚFSDPv2r    Úis_fsdp_xla_v2_enableds       €€€€r   Úauto_wrapper_callablez2wrap_model_xla_fsdp.<locals>.auto_wrapper_callablet   s'   ø€ Ù%;™ÀˆJÙÑ/°Ó2ÐD°TÒD¸VÑDÐDr   c                 óì   •— ddl m} d }t        | t        j                  «      r| }n.t        | t
        «      r| d   }nt        | |«      r| j                  }|€t        d«      ‚‰j                  ||d«       y )Nr   )ÚCausalLMOutputWithPastr   zASomething went wrong, the output of the model shouldn't be `None`)r   NN)	Úmodeling_outputsr5   r   ÚtorchÚTensorÚtupleÚlogitsÚ
ValueErrorÚmark_sharding)ÚoutputÚmeshr5   Úreal_outputr   s       €r   Úshard_outputz)wrap_model_xla_fsdp.<locals>.shard_output{   sj   ø€ ÝAàˆKÜ˜&¤%§,¡,Ô/Ø$‘Ü˜F¤EÔ*Ø$ Q™i‘Ü˜FÐ$:Ô;Ø$Ÿm™m�àÐ"Ü Ð!dÓeÐeØ×Ñ˜[¨$Ð0DÕEr   )r@   Úauto_wrap_policyr3   )rA   r3   c                 óP   •—  | j                   di |¤Ž}|r‰j                  «        |S )Nr+   )ÚstepÚ	mark_step)Ú	optimizerÚbarrierÚoptimizer_argsÚlossÚxms       €r   Úpatched_optimizer_stepz3wrap_model_xla_fsdp.<locals>.patched_optimizer_stepš   s'   ø€ Øˆy�~‰~Ñ/ Ñ/ˆÙØ�L‰LŒNØˆr   )Útorch_xla.core.xla_modelÚcoreÚ	xla_modelr   r   r   Útrainer_pt_utilsr   Útorch_xla.distributed.fsdpr   r    Útorch_xla.distributed.fsdp.wrapr!   r"   Ú7torch_xla.experimental.spmd_fully_sharded_data_parallelr#   ÚImportErrorÚgetattrÚfsdp_configÚgetÚ	functoolsÚpartialÚsetÚ	ExceptionÚaddÚxla_fsdp_configÚconfigÚ	use_cacheÚloggerÚwarning_onceÚoptimizer_step)Úmodelr-   r2   r   r!   r"   rA   r3   Ú%default_transformer_cls_names_to_wrapÚ"fsdp_transformer_layer_cls_to_wrapÚtransformer_cls_to_wrapÚlayer_classÚtransformer_clsÚfsdp_kwargsr@   rJ   r0   r1   r    rI   r   s     `             @@@@@r   Úwrap_model_xla_fsdprh   .   s×  ý€ ÷ *Ð)ß+Ð+å=ðhÝRÝ@÷	
ñ
 "õð ÐØ ÐÜ,3°EÐ;NÐPTÓ,UÐ)Ø)-×)9Ñ)9×)=Ñ)=Ø'Ð)Nó*Ð&ð ×ÑÐ(Ñ)¨AÒ-Ü$×,Ñ,Ø'¸×8HÑ8HÐIYÑ8Zô
Ñð 
,Ð	7Ü"%£%ÐØ=ò 	=ˆKÙ8¸ÀÓLˆOØÐ&ÜÐ bÓcÐcà'×+Ñ+¨OÕ<ð	=ô %×,Ñ,Ø(à"9ô
Ðð ×&Ñ&€KØ×ÑÐ,Ò-Ø�<‰<×!Ò!Ü×ÑØjôð &+ˆE�L‰LÔ"÷	Eñ
 ô	Fñ ØØ%Ø-Ø"7ô	
‰ñ Øð
à-Ø"7ñ
ð ñ	
ˆð 38Èõ ð /€BÔà€Løôi ò hÜÐfÓgÐgðhús    F ÆFc           	      ó¤  — ddl mc m} |�|n|j                  }t        j                  d|› �«       |j                  «        |j                  d¬«      rKt        j                  |d¬«       t        j                  |t        j                  j                  |d«      «       t        f}|j                  d	«       |�r`| j!                  «       | j#                  «       d
œ}t        j                  j                  |d|j$                  › d|j&                  › dt(        › �«      }	|j                  ||	d¬«       |j                  d«       |j*                  �râddlm}
  |
t        j                  j                  |d«      dt(        › �d¬«      \  }}| j0                  j0                  } |j3                  | «      }t5        ||«      r|j7                  ||¬«       �nat        j                  d«       |j                  |t        j                  j                  |t(        «      «       �nt5        | |«      sÏt5        |j3                  | «      |«      rK|j3                  | «      j7                  ||j*                  |j9                  | j!                  «       «      ¬«       n¤t        j                  d«       |j9                  | j!                  «       «      }|j                  |t        j                  j                  |t(        «      «       n;| j7                  ||j*                  |j9                  | j!                  «       «      ¬«       |�|j*                  r|j7                  |«       yyy)a”  
    Saves a model checkpoint on TPU/XLA devices.

    Handles FSDP v1 sharded checkpoints (with consolidation on master), as well as
    standard XLA model saving via `save_pretrained` or `xm.save`.

    Args:
        model (`torch.nn.Module`): The model to save.
        args (`TrainingArguments`): The training arguments.
        accelerator (`Accelerator`): The accelerator instance.
        processing_class: The processing class (tokenizer/processor) to save alongside the model.
        is_fsdp_xla_v1_enabled (`bool`): Whether FSDP XLA v1 is enabled.
        output_dir (`str`, *optional*): The directory to save to. Defaults to `args.output_dir`.
    r   NzSaving model checkpoint to F)ÚlocalT)Úexist_okztraining_args.binÚsaving_checkpoint)ra   Úshard_metadataÚrankz-of-ú-)Úmaster_onlyÚsave_full_checkpoints)Ú%consolidate_sharded_model_checkpointsÚ zrank*-of-*-)Úckpt_prefixÚckpt_suffixÚ
save_model)Ú
state_dictzETrainer.model is not a `PreTrainedModel`, only saving its state dict.)Úis_main_processrw   )rK   rL   rM   Ú
output_dirr^   ÚinforD   Úis_master_ordinalÚosÚmakedirsr7   ÚsaveÚpathÚjoinr   Ú
rendezvousrw   Úget_shard_metadataÚprocess_indexÚ
world_sizer   Úshould_saverO   rr   ÚmoduleÚunwrap_modelr   Úsave_pretrainedÚ_maybe_convert_to_cpu)ra   r-   ÚacceleratorÚprocessing_classÚis_fsdp_xla_v1_enabledry   rI   Úsupported_classesÚckptÚ	ckpt_pathrr   Úfull_state_dictÚ_Úunwrapped_modelrw   s                  r   Úsave_tpu_checkpointr“   ¥   s«  € ÷ *Ð)à)Ð5‘¸4¿?¹?€Jä
‡K�KÐ-¨j¨\Ð:Ô;Ø‡L�L„Nà	×Ñ %ÐÔ(Ü
�‰�J¨Õ.Ü�
‰
�4œŸ™Ÿ™ jÐ2EÓFÔGô (Ð)ÐØ‡M�MÐ%Ô&Úà×%Ñ%Ó'Ø#×6Ñ6Ó8ñ
ˆô —G‘G—L‘L ¨t°D×4FÑ4FÐ3GÀtÈDÏOÉOÐK\Ð\]Ô^jÐ]kÐ-lÓmˆ	à
�‰��i¨UˆÔ3à
�‰Ð-Ô.à×ÓÝXá!FÜŸG™GŸL™L¨°RÓ8Ø)¬,¨Ð8Ø ô"ÑˆO˜Qð
 —L‘L×'Ñ'ˆEØ)×6Ñ6°uÓ=ˆOÜ˜/Ð+<Ô=Ø×/Ñ/°
ÀÐ/ÖWä—‘ÐcÔdØ—‘˜¬¯©¯©°jÄ,Ó)OÖPÜ˜Ð0Ô1Ü�k×.Ñ.¨uÓ5Ð7HÔIØ×$Ñ$ UÓ+×;Ñ;ØØ $× 0Ñ 0Ø×3Ñ3°E×4DÑ4DÓ4FÓGð <õ ô �K‰KÐ_Ô`Ø×1Ñ1°%×2BÑ2BÓ2DÓEˆJØ�G‰G�J¤§¡§¡¨Z¼Ó FÕGà×ÑØØ ×,Ñ,Ø×/Ñ/°×0@Ñ0@Ó0BÓCð 	ô 	
ð
 Ð#¨×(8Ò(8Ø×(Ñ(¨Õ4ð )9Ð#r   r*   )rV   r|   r7   Útorch.utils.datar   Úutilsr   r   r   r   Ú
get_loggerÚ__name__r^   r   rh   r“   r+   r   r   ú<module>r˜      sF   ðó Û 	ã Ý 'ç QÓ Qð 
ˆ×	Ñ	˜HÓ	%€ð Jó ò&tônJ5r   