Llama-3.2-3B-Instruct โ NAS-Pruned (PubMedQA assessment)
NAS depth+width pruned variant of meta-llama/Llama-3.2-3B-Instruct, targeting 2.6B
parameters (19% reduction from ~3.21B) โ the pruning step of a medical-domain
(PubMedQA) compression pipeline built for the DLI course capstone assessment.
Base model: Decoder-only Transformer, 28 layers, hidden_size=3072, ffn_hidden_size=8192, 24 attention heads, 8 GQA key-value groups (~3.21B total params).
Compression Method
torchrun --nproc_per_node 1 prune_minitron.py \
--hf_model_name_or_path checkpoints/Llama-3.2-3B-Instruct \
--calib_dataset_name wikitext \
--prune_target_params 2.6e9 \
--calib_num_samples 256 \
--top_k 3 \
--prune_score_func "mmlu_1pct" \
--trust_remote_code \
--output_hf_path pruning_assessment/Llama-3.2-3B-Instruct-pruned
Not yet distilled โ see tocsa/llama-3.2-3b-pruned-distilled-fp8 for the
distilled + FP8-quantized descendant.
Evaluation
Evaluated 5-shot with lm_eval on mmlu_anatomy (MMLU anatomy subtask, chosen as
representative for the PubMedQA medical domain):
| Checkpoint | Target | Score |
|---|---|---|
Baseline Llama-3.2-3B-Instruct |
> 0.59 | fill in |
| This pruned checkpoint | > 0.35 | fill in |
Framework
Produced with NVIDIA Model-Optimizer / Megatron-Bridge, as the capstone assessment of the DLI course "The Art of Compressing LLMs: Pruning, Distillation, and Quantization Demystified." Domain dataset: PubMedQA.
License
Inherits the license terms of the base model, meta-llama/Llama-3.2-3B-Instruct
(Llama 3.2 Community License). Redistribution must comply with that license.
- Downloads last month
- 3