
    ^j*                    h   d dl mZ d dlZd dlZd dlmZmZ ddlmZ ddl	m
Z
mZmZmZ ddlmZ erd dlmZ dd	lmZ  e
       rd dlZ e
       r ed
      rd dlmZ d dlmZmZ  ej4                  e      ZddZddZddZddZ ddZ!	 	 	 	 	 	 ddZ"	 	 	 	 	 	 ddZ#ddZ$	 	 	 	 	 	 d dZ%d Z&d Z'y)!    )annotationsN)TYPE_CHECKINGAny    replace_layer_number_by_wildcard)is_torch_availableis_torch_greater_or_equallogging	strtobool)QuantizationMethod   )DistributedConfig2.6)fully_shard)CPUOffloadPolicyMixedPrecisionPolicyc                 L   t               syt        j                  j                         xrz t        j                  j	                         xrZ t        t        j                  j                  dd            dk(  xr, t        t        j                  j                  dd            dk(  S )uM   Check if FSDP is active via Accelerate (env var based) — covers FSDP1 only.FACCELERATE_USE_FSDPFalser   FSDP_CPU_RAM_EFFICIENT_LOADING)	r	   torchdistributedis_availableis_initializedr   osenvironget     h/var/www/ramen.bs-engineer-server.com/venv/lib/python3.12/site-packages/transformers/distributed/fsdp.pyis_fsdp_enabledr"   (   s     	&&( 	V,,.	Vbjjnn%:GDEJ	V bjjnn%EwOPTUU	r    c                    t               syt        j                  j                         syt	        | dd      ry	 ddlm} t        | |      S # t        $ r Y yw xY w)z.Check if a module is managed by FSDP (1 or 2).F_is_fsdp_managed_moduleTr   )FullyShardedDataParallel)	r	   r   r   r   getattrtorch.distributed.fsdpr%   ImportError
isinstance)moduler%   s     r!   is_fsdp_managed_moduler+   5   s^    ))+ v0%8C f677  s   A 	AAc                    | i S i }| j                   rt               |d<   | j                  r-t        t        j
                  t        j                  d      |d<   |S )zMBuild ``fully_shard`` policy kwargs from ``DistributedConfig`` runtime flags.Noffload_policy)param_dtypereduce_dtypeoutput_dtype	mp_policy)fsdp_cpu_offloadr   fsdp_mixed_precisionr   r   bfloat16float32)distributed_configfsdp_policy_kwargss     r!   _get_fsdp_policy_kwargsr8   G   s]    !	**/?/A+,..*>+
;'
 r    c                    d }d }t        | d      r| j                         }t        | d      r| j                         }||fS )Nget_input_embeddingsget_output_embeddings)hasattrr:   r;   )modelinput_embedoutput_heads      r!   _get_input_output_embeddingsr@   X   sI    KKu,-002u-.113##r    c                *  
 t        |       dk7  ryt        |      \  }}||fD ch c]  }||	 c}
g g }}| D ]'  \  }}|j                  |       |j                  |       ) t        d |D              }t        
fd|D              }	|xr |	S c c}w )Nr   Fc              3  L   K   | ]  }|d k(  xs |j                  d        yw)normz.normN)endswith).0names     r!   	<genexpr>z(is_norm_and_head_pair.<locals>.<genexpr>m   s%     TdA4==+AATs   "$c              3  &   K   | ]  }|v  
 y wNr   )rE   r*   head_moduless     r!   rG   z(is_norm_and_head_pair.<locals>.<genexpr>n   s     GV&L0Gs   )lenr@   appendany)no_reshard_targetsr=   r>   r?   r*   namesmodulesrF   has_final_normhas_output_headrJ   s             @r!   is_norm_and_head_pairrS   b   s    
!#;EBK*5{)CZvvGYFZL7E* fTv TeTTNGwGGO-o- [s
   BBc                   t        |dd      xs i }|s| S t        |      \  }}|j                         D ci c]  \  }}||
 }}}|j                  |      }|j                  |      }	||	| S | j	                         }
|
j                  |d       | j                  |	      dk(  r|
j                  |	d       d|
|<   |
S c c}}w )a  
    Rewrite the plan so tied embed/lm_head weights are wrapped once.
    Example:
        {"model.embed_tokens": "free_full_weight",
        "model.layers.*": "free_full_weight",
        "model.norm": "keep_full_weight",
        "lm_head": "keep_full_weight"}
    ->
        {"model.layers.*": "free_full_weight",
        "model.norm": "keep_full_weight",
        "model.embed_tokens": "keep_full_weight"}
    all_tied_weights_keysNkeep_full_weight)r&   r@   named_modulesr   copypop)	fsdp_planr=   	tied_keysr>   r?   rF   r*   name_by_moduleembed_modulehead_moduleadapted_plans              r!    _resolve_tied_embed_lm_head_planr`   r   s      6=CI;EBK7<7J7J7LM|tVfdlMNM!%%k2L $$[1K{2>>#L\4(}}[!%77d+%7\" Ns   B>c                    g }g }| j                         D ]J  \  }}||v r|n
t        |      }||v s||   dk(  r|j                  ||f       8|j                  ||f       L ||fS )zUExpand plan keys into reshard and no-reshard ``(module_name, module)`` shard targets.rV   )rW   r   rL   )r=   rZ   reshard_targetsrN   module_namer*   plan_keys          r!   expand_fsdp_planre      s    
 46O68$224 >V"-":;@`al@my "&88"));*?@&&V'<=> ...r    c                *   |syt         j                  |       }i }i }|j                         D ].  \  }|dvr||<   |vst        fd|D              r*||<   0 |rt        j                  d|        |rt        j                  d|        yy)zs
    Verify the FSDP plan of the model, log a warning if plan keys were not applied or strategies are invalid.
    N>   free_full_weightrV   c              3  :   K   | ]  }t        |      k(    y wrI   r   )rE   rF   keys     r!   rG   z#verify_fsdp_plan.<locals>.<genexpr>   s     /vbf0PQU0VZ]0]/vs   z4The following FSDP entries have unknown strategies: z9The following FSDP rules were not applied to any module: )dictfromkeysitemsrM   loggerwarning)module_namesrZ   name_lookupunused_rulesinvalid_strategiesstrategyri   s         @r!   verify_fsdp_planrt      s     ---K#%L)+"* )XCC&.s##C/vju/v,v (L	) MN`MabcRS_R`ab r    c                P   t               st        d      t        d      st        d      t	        t        | dd      xs i       }|s!t        t        |       j                   d      t        | j                  dd      }t        |      }t        ||       }t        | |      \  }}|D ]-  \  }}	t        |	f|dd	| t        j                  d
| d       / t!        ||       rYg g }}
|D ]'  \  }}	|
j#                  |       |j#                  |	       ) t        |f|dd	| t        j                  d|
 d       n2|D ]-  \  }}	t        |	f|dd	| t        j                  d
| d       / t        | fd|i| t        j%                  dt'        |       d       d| _        | S )z/
    Apply FSDP2 (fully_shard) to a model.
    z$PyTorch is required for FSDP supportr   zFSDP2 requires torch>=2.6
_fsdp_planNzr does not have a FSDP2 plan declared. Set `base_model_fsdp_plan` on the config and `_fsdp_plan` on the head class.r6   T)meshreshard_after_forwardzApplied fully_shard to z (reshard=True)FzGrouped tail z (reshard=False)rw   z'FSDP2 applied to model via _fsdp_plan: z entries)r	   r(   r
   OSErrorrj   r&   
ValueErrortype__name__configr8   r`   re   r   rm   debugrS   rL   inforK   r$   )r=   	fsdp_meshrZ   r6   r7   adapted_fsdp_planrb   rN   rc   r*   rO   rP   rF   s                r!   !apply_fully_sharded_data_parallelr      s    @AA$U+122WUL$7=2>IE{##$ %W W
 	

 !/CTJ01CD8EJ*:5BS*T'O'. MVF]$]J\].{m?KLM /7Rw. 	#LD&LLNN6"	# 	G_)5_L^_}UG+;<=. 	KLD&bYebOabLL24&8HIJ	K
 <I<);<
KK9#i.9IRS %)E! Lr    c                 n    ddl m}  dt        t        j                  |       j
                        v rddiS i S )z
    Returns checkpoint kwargs for FSDP model saving.

    Checks if the `adapter_only` parameter is supported by `save_fsdp_model` from accelerate
    and returns the appropriate kwargs.
    r   save_fsdp_modeladapter_onlyT)accelerate.utilsr   listinspect	signature
parametersr   s    r!   get_fsdp_ckpt_kwargsr      s5     1g//@KKLL%%	r    c                   ddl m} ddlm} t	        | j
                  |      r! ||       |j                  j                  _        t        | dd      t        j                  k(  rq| j                  j                  j                  j                  rF|j                  j                  j!                  | j                  j                  j                  d       yyy)aG  
    Updates the FSDP plugin for PEFT LoRA/QLoRA compatibility.

    When using FSDP with PEFT LoRA, the auto wrap policy needs to be updated to additionally wrap
    LoRA trainable layers separately. When using FSDP with QLoRA, the mixed precision policy needs
    to be updated to use the quantization storage data type.
    r   )
PeftConfig)fsdp_auto_wrap_policyquantization_methodNT)override)peftr   peft.utils.otherr   r)   active_peft_configstatefsdp_pluginauto_wrap_policyr&   r   BITS_AND_BYTEShf_quantizerquantization_configbnb_4bit_quant_storageis_floating_pointset_mixed_precision)r=   acceleratorr   r   s       r!   update_fsdp_plugin_peftr     s      6%**J79Nu9U%%6,d37I7X7XX22II[[%%9922IITX 	: 	
 \ 	Yr    )returnbool)r*   	nn.Moduler   r   )r6   zDistributedConfig | Noner   zdict[str, Any])r=   r   r   z)tuple[nn.Module | None, nn.Module | None])rN   zlist[tuple[str, nn.Module]]r=   r   r   r   )rZ   dict[str, str]r=   r   r   r   )r=   r   rZ   r   r   z?tuple[list[tuple[str, nn.Module]], list[tuple[str, nn.Module]]])ro   z	list[str]rZ   zdict[str, str] | Noner   None)r=   r   r   z(torch.distributed.device_mesh.DeviceMeshr   r   )(
__future__r   r   r   typingr   r   integrations.tensor_parallelr   utilsr	   r
   r   r   utils.quantization_configr   torch.nnnnconfiguration_utilsr   r   "torch.distributed._composable.fsdpr   r'   r   r   
get_loggerr|   rm   r"   r+   r8   r@   rS   r`   re   rt   r   r   r   r   r    r!   <module>r      s    #  	 % K U U : 65e<>M			H	%
8$"$. ### #L/// E/&c.55!I55t
r    