o
    �Tj×  ã                   @  s|   d Z ddlmZ ddlZddlZddlmZ ej de	ee
ƒ ¡ jd d ƒ¡ ddlmZ dd
d„Zedkr<eƒ  dS dS )u  QLoRA fine-tune of Gemma 4 (E2B/E4B) on the SFT set.

    python scripts/train_qlora.py --config config/config.yaml       # (GPU required)

Primary path uses **Unsloth** (day-0 Gemma 4 support, ~2x faster / ~50% less VRAM, and it handles
the E-series multimodal wrapper + known gotchas for you). Only the LoRA adapter is trained; the
4-bit base weights stay frozen â€” that's what keeps VRAM/energy low. VRAM: E2B ~8-10GB, E4B ~17GB.

The adapter is written to train.output_adapter and merged/quantized by merge_and_quantize.py.

Gemma-4 QLoRA notes:
  - Gemma 4 supports a real system role, so our system/user/assistant SFT messages apply cleanly.
  - Training loss in the 13-15 range is normal for the E-series (multimodal) models.
  - Use bf16 (not fp16) to avoid overflow on the E-series.
é    )ÚannotationsN)ÚPathé   Úsrc)Úload_configÚreturnÚNonec                    sÎ  t  ¡ } | jddd� |  ¡ }t|jƒ}ddlm} ddlm	}m
} ddlm} ddlm} | d	¡}t| d
¡ƒ}	t| d¡ƒ}
t| dd¡ƒ}|j||ddd�\}‰ |ˆ dd�‰ |j|t| dd¡ƒt| dd¡ƒt| dd¡ƒd| d¡ddd�}|d|	dd �}‡ fd!d"„}|j||jd#�}|tt|
ƒjd$ ƒt| d%d&¡ƒt| d'd&¡ƒt| d(d)¡ƒt| d*d+¡ƒ|ddd,d-d.d/dd0d1�}||ˆ ||d2�}| ¡  | |
¡ ˆ  |
¡ td3|
› �ƒ td4|jƒ d S )5Nz--configzconfig/config.yaml)Údefaultr   )Úload_dataset)Ú	SFTConfigÚ
SFTTrainer)Ú	FastModel)Úget_chat_templatezmodel.base_idzdata.sft_trainztrain.output_adapterztrain.max_seq_leni   TF)Z
model_nameÚmax_seq_lengthZload_in_4bitZfull_finetuningzgemma-4)Zchat_templateztrain.lora_ré   ztrain.lora_alphaé    ztrain.lora_dropoutgš™™™™™©?Znoneztrain.target_modulesÚunslothiO  )ÚrZ
lora_alphaZlora_dropoutZbiasZtarget_modulesZuse_gradient_checkpointingZrandom_stateZjsonÚtrain)Z
data_filesÚsplitc                   s   dˆ j | d ddd�iS )NÚtextÚmessagesF)ÚtokenizeZadd_generation_prompt)Zapply_chat_template)Zexample©Ú	tokenizer© úscripts/train_qlora.pyÚformat_chatC   s   
ÿÿzmain.<locals>.format_chat)Zremove_columnsÚtrainerztrain.epochsé   ztrain.batch_sizeztrain.grad_accumé   ztrain.lrg-Cëâ6*?é   ZepochZcosineg¸…ëQ¸ž?r   )Z
output_dirZnum_train_epochsZper_device_train_batch_sizeZgradient_accumulation_stepsZlearning_rater   Zbf16ZpackingZlogging_stepsZsave_strategyZlr_scheduler_typeZwarmup_ratioZ	report_toZdataset_text_field)Úmodelr   ÚargsZtrain_datasetzSaved LoRA adapter -> z3Next: python scripts/merge_and_quantize.py --config)ÚargparseÚArgumentParserÚadd_argumentÚ
parse_argsr   ZconfigZdatasetsr
   Ztrlr   r   r   r   Zunsloth.chat_templatesr   ÚgetÚstrÚresolveÚintZfrom_pretrainedZget_peft_modelÚfloatÚmapZcolumn_namesr   Úparentr   Zsave_pretrainedÚprint)Zapr#   Zcfgr
   r   r   r   r   Zbase_idZ
train_fileZoutput_adapterZmax_seq_lenr"   Zdatasetr   Zsft_cfgr   r   r   r   Úmain   sl   


üøò

r0   Ú__main__)r   r   )Ú__doc__Z
__future__r   r$   ÚsysZpathlibr   ÚpathÚinsertr)   Ú__file__r*   ÚparentsZagent.settingsr   r0   Ú__name__r   r   r   r   Ú<module>   s    $
Y
ÿ