o
    5�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ddd„Zd dd„Zd!dd„Zd"dd„Zd#dd„ZedkrUeƒ  dS dS )$a‹  Merge the LoRA adapter into Gemma 4 and quantize to int4 for serving.

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

Steps:
  1. Load the 4-bit base + LoRA adapter via Unsloth and merge to a standalone 16-bit model.
  2. Quantize per model.quantization:
       - w4a16 -> compressed-tensors int4 via llm-compressor  (recommended: day-0 vLLM support,
                  same format as Gemma 4's official QAT weights)
       - gguf  -> Q4_0 GGUF via Unsloth  (for Ollama / llama.cpp; Q4_0 matches Gemma 4 QAT)
       - awq   -> AutoAWQ int4
  3. Write to model.quantized_dir.

int4 serving is the biggest inference-time energy win: an E4B model in int4 needs only a few GB of
weights. Tip: if you don't need custom fine-tuned behavior, you can skip training entirely and
serve Google's official QAT checkpoint (google/gemma-4-E4B-it-qat-*), which is already w4a16.
é    )ÚannotationsN)ÚPathé   Úsrc)Úload_configÚbase_idÚstrÚadapter_dirÚmax_seq_lenÚintc                 C  s0   ddl m} tdƒ |j||dd�\}}||fS )zKLoad base + adapter via Unsloth (returns model, tokenizer ready to export).r   )Ú	FastModelz(Loading base + LoRA adapter (Unsloth)...F)Z
model_nameÚmax_seq_lengthZload_in_4bitN)Zunslothr   ÚprintÚfrom_pretrained)r   r	   r
   r   ÚmodelÚ	tokenizer© r   úscripts/merge_and_quantize.pyÚload_merged   s   
ýr   Úout_dirc                 C  s2   t dƒ | j||dd� t d|› �ƒ t dƒ d S )Nz=Exporting merged model to GGUF (Q4_0) for Ollama/llama.cpp...Zq4_0)Zquantization_methodzGGUF -> zRServe with Ollama: create a Modelfile `FROM ./<model>.gguf` then `ollama create`. )r   Zsave_pretrained_gguf)r   r   r   r   r   r   Úexport_gguf*   s   r   Ú
merged_dirc                 C  s*   t dƒ | j||dd� t d|› �ƒ |S )Nz"Merging LoRA into base (16-bit)...Zmerged_16bit)Zsave_methodzMerged 16-bit -> )r   Zsave_pretrained_merged)r   r   r   r   r   r   Úexport_merged_16bit1   s   r   c                 C  sX   ddl m} ddlm} tdƒ |ddg d¢d�}|| d	|||d
d� td|› �ƒ dS )zIcompressed-tensors W4A16 via llm-compressor (GPTQ-style), served by vLLM.r   )ÚGPTQModifier)Úoneshotz>Quantizing to W4A16 (compressed-tensors) via llm-compressor...ZLinearZW4A16)Zlm_headzre:.*vision_tower.*zre:.*audio_tower.*zre:.*multi_modal.*)ÚtargetsZschemeÚignoreZopen_platypusé   )r   ZdatasetÚrecipeZ
output_dirr   Znum_calibration_sampleszW4A16 model -> N)Z$llmcompressor.modifiers.quantizationr   Zllmcompressor.transformersr   r   )r   r   r
   r   r   r   r   r   r   Úquantize_w4a168   s"   üúr   c                 C  sr   ddl m} ddlm} tdƒ | | ¡}| | ¡}|j|ddddd	œd
� | |¡ | |¡ td|› �ƒ d S )Nr   )ÚAutoAWQForCausalLM)ÚAutoTokenizerzQuantizing with AWQ (int4)...Té€   é   ZGEMM)Z
zero_pointZq_group_sizeZw_bitÚversion)Zquant_configzAWQ int4 model -> )	Úawqr    Ztransformersr!   r   r   ZquantizeZsave_quantizedZsave_pretrained)r   r   r    r!   r   r   r   r   r   Úquantize_awqO   s   



r&   ÚreturnÚNonec                  C  s$  t  ¡ } | jddd� |  ¡ }t|jƒ}| d¡}t| d¡ƒ}t| d¡ƒ}t| dd¡ƒ 	¡ }t
| d	d
¡ƒ}t|||ƒ\}}	|dkrMt||	|ƒ d S t| d¡jd ƒ}
t||	|
ƒ |dkrxt|
||ƒ tdƒ td|› d|› d�ƒ d S |dkr‹t|
|ƒ td|› d�ƒ d S td|› �ƒ‚)Nz--configzconfig/config.yaml)Údefaultzmodel.base_idztrain.output_adapterzmodel.quantized_dirzmodel.quantizationZw4a16zmodel.max_model_leni   Zggufzmodel-merged-16bitz
Serve with vLLM:z7  python -m vllm.entrypoints.openai.api_server --model z --max-model-len z --port 8001r%   z
Serve with vLLM: ... --model z --quantization awq --port 8001zUnknown quantization method: )ÚargparseÚArgumentParserÚadd_argumentÚ
parse_argsr   ZconfigÚgetr   ÚresolveÚlowerr   r   r   Úparentr   r   r   r&   Ú
SystemExit)ZapÚargsZcfgr   r	   r   Úmethodr
   r   r   r   r   r   r   Úmain\   s2   


ÿ
r5   Ú__main__)r   r   r	   r   r
   r   )r   r   )r   r   )r   r   r   r   r
   r   )r   r   r   r   )r'   r(   )Ú__doc__Z
__future__r   r*   ÚsysZpathlibr   ÚpathÚinsertr   Ú__file__r/   ÚparentsZagent.settingsr   r   r   r   r   r&   r5   Ú__name__r   r   r   r   Ú<module>   s     $





!
ÿ