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
Safetensors
Model size
3B params
Tensor type
BF16
ยท
Inference Providers NEW
This model isn't deployed by any Inference Provider. ๐Ÿ™‹ Ask for provider support