Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 4 additions & 1 deletion .gitignore
Original file line number Diff line number Diff line change
Expand Up @@ -13,4 +13,7 @@ site/
.cache/
**mase/
**fast-hadamard-transform/
**calib/
**calib/
.venv/
.pytest_cache/
*.vcd
2 changes: 1 addition & 1 deletion quant_eval/cli/QUANT_HF_SERVE.md
Original file line number Diff line number Diff line change
Expand Up @@ -19,7 +19,7 @@ python -m quant_eval.cli.quant_hf_serve \
```bash
python -m quant_eval.cli.quant_hf_serve \
--model_name Qwen/Qwen3-8B \
--quant_config quant_eval/configs/llama_mxint4.toml \
--quant_config quant_eval/configs/llama_mxint8.toml \
--prefill_attn_width 4 --prefill_ffn_width 4 \
--decode_attn_width 8 --decode_ffn_width 8 \
--host 127.0.0.1 --port 8915
Expand Down
2 changes: 1 addition & 1 deletion quant_eval/cli/eval_dllm.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,7 +17,7 @@

python -m quant_eval.cli.eval_dllm \\
--model_name Efficient-Large-Model/Fast_dLLM_v2_1.5B \\
--quant_config quant_eval/configs/llama_mxint4.toml \\
--quant_config quant_eval/configs/llama_mxint8.toml \\
--tasks gsm8k
"""

Expand Down
4 changes: 2 additions & 2 deletions quant_eval/cli/eval_evalplus.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,7 +14,7 @@

python -m quant_eval.cli.eval_evalplus \\
--model_name unsloth/Llama-3.2-1B \\
--quant_config quant_eval/configs/llama_mxint4.toml \\
--quant_config quant_eval/configs/llama_mxint8.toml \\
--dataset humaneval \\
--greedy \\
--evalplus_output_dir logs/evalplus
Expand Down Expand Up @@ -47,7 +47,7 @@ def main(
dataset: str = "humaneval",
device_id: str = "cuda:0",
dtype: str = "bfloat16",
quant_config: Union[str, None] = "quant_eval/configs/llama_mxint4.toml",
quant_config: Union[str, None] = "quant_eval/configs/llama_mxint8.toml",
model_parallel: bool = False,
batch_size: int = 1,
greedy: bool = False,
Expand Down
2 changes: 1 addition & 1 deletion quant_eval/cli/eval_llada.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,7 +16,7 @@
python -m quant_eval.cli.eval_llada \\
--tasks gsm8k --num_fewshot 0 \\
--model llada_dist \\
--model_args model_path='GSAI-ML/LLaDA-8B-Instruct',gen_length=256,steps=256,block_length=32,use_cache=True,quant_config='quant_eval/configs/llama_mxint4.toml'
--model_args model_path='GSAI-ML/LLaDA-8B-Instruct',gen_length=256,steps=256,block_length=32,use_cache=True,quant_config='quant_eval/configs/llama_mxint8.toml'
"""

# Import to trigger @register_model("llada_dist")
Expand Down
4 changes: 2 additions & 2 deletions quant_eval/cli/eval_lm.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,7 @@

python -m quant_eval.cli.eval_lm \\
--model_name unsloth/Llama-3.2-1B \\
--quant_config quant_eval/configs/llama_mxint4.toml \\
--quant_config quant_eval/configs/llama_mxint8.toml \\
--tasks arc_easy,hellaswag,winogrande \\
--limit 500
"""
Expand Down Expand Up @@ -40,7 +40,7 @@ def main(
tasks: Union[str, list[str]] = "wikitext",
device_id: str = "cuda:0",
dtype: str = "bfloat16",
quant_config: Union[str, None] = "quant_eval/configs/llama_mxint4.toml",
quant_config: Union[str, None] = "quant_eval/configs/llama_mxint8.toml",
model_parallel: bool = False,
seqlen: int = 2048,
batch_size: Union[int, str] = 64,
Expand Down
2 changes: 1 addition & 1 deletion quant_eval/cli/eval_osworld.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,7 +18,7 @@
python -m quant_eval.cli.eval_osworld \\
--model_name Qwen/Qwen2.5-7B-Instruct \\
--osworld_path quant_eval/benchmarks/OSWorld \\
--quant_config quant_eval/configs/llama_mxint4.toml \\
--quant_config quant_eval/configs/llama_mxint8.toml \\
--domain chrome --max_steps 15
"""

Expand Down
4 changes: 2 additions & 2 deletions quant_eval/cli/eval_phase_bfcl.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,7 +23,7 @@

python -m quant_eval.cli.eval_phase_bfcl \\
--model_name Qwen/Qwen2.5-1.5B \\
--quant_config quant_eval/configs/llama_mxint4.toml \\
--quant_config quant_eval/configs/llama_mxint8.toml \\
--prefill_attn_width 4 --prefill_ffn_width 4 \\
--decode_attn_width 8 --decode_ffn_width 8 \\
--bfcl_test_categories web_search_base \\
Expand Down Expand Up @@ -599,7 +599,7 @@ def main(
model_name: str = "Qwen/Qwen3-8B-FC",
device_id: str = "cuda:0",
dtype: str = "bfloat16",
quant_config: str = "quant_eval/configs/llama_mxint4.toml",
quant_config: str = "quant_eval/configs/llama_mxint8.toml",
model_parallel: bool = False,
# ── BFCL settings ──────────────────────────────────────────────────────
bfcl_test_categories: Union[list[str], None] = None,
Expand Down
4 changes: 2 additions & 2 deletions quant_eval/cli/eval_phase_lm.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,7 +15,7 @@

python -m quant_eval.cli.eval_phase_lm \\
--model_name Qwen/Qwen2.5-1.5B \\
--quant_config quant_eval/configs/llama_mxint4.toml \\
--quant_config quant_eval/configs/llama_mxint8.toml \\
--prefill_attn_width 4 --prefill_ffn_width 4 \\
--decode_attn_width 8 --decode_ffn_width 8 \\
--tasks gsm8k --limit 200
Expand Down Expand Up @@ -49,7 +49,7 @@ def main(
tasks: Union[str, list[str]] = "wikitext",
device_id: str = "cuda:0",
dtype: str = "bfloat16",
quant_config: str = "quant_eval/configs/llama_mxint4.toml",
quant_config: str = "quant_eval/configs/llama_mxint8.toml",
model_parallel: bool = False,
seqlen: int = 2048,
batch_size: Union[int, str] = 64,
Expand Down
2 changes: 1 addition & 1 deletion quant_eval/cli/eval_ppl.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,7 +13,7 @@

python -m quant_eval.cli.eval_ppl \\
--model_name unsloth/Llama-3.2-1B \\
--quant_config quant_eval/configs/llama_mxint4.toml
--quant_config quant_eval/configs/llama_mxint8.toml
"""

from typing import Union
Expand Down
10 changes: 5 additions & 5 deletions quant_eval/cli/quant_hf_serve.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,7 +24,7 @@
Example:
python -m quant_eval.cli.quant_hf_serve \\
--model_name Qwen/Qwen3-8B-FC \\
--quant_config quant_eval/configs/llama_mxint4.toml \\
--quant_config quant_eval/configs/llama_mxint8.toml \\
--prefill_attn_width 4 --prefill_ffn_width 4 \\
--decode_attn_width 8 --decode_ffn_width 8 \\
--host 127.0.0.1 --port 8915
Expand Down Expand Up @@ -356,15 +356,15 @@ def main(
if model_parallel:
model = AutoModelForCausalLM.from_pretrained(
model_name,
torch_dtype=torch_dtype,
dtype=torch_dtype,
device_map="auto",
trust_remote_code=True,
)
print(f"Device map: {model.hf_device_map}")
else:
model = AutoModelForCausalLM.from_pretrained(
model_name,
torch_dtype=torch_dtype,
dtype=torch_dtype,
trust_remote_code=True,
).to(device_id)
model.eval()
Expand Down Expand Up @@ -419,7 +419,7 @@ def main(
if model_parallel:
model = AutoModelForCausalLM.from_pretrained(
model_name,
torch_dtype=torch_dtype,
dtype=torch_dtype,
device_map="auto",
trust_remote_code=True,
)
Expand All @@ -428,7 +428,7 @@ def main(
else:
model = AutoModelForCausalLM.from_pretrained(
model_name,
torch_dtype=torch_dtype,
dtype=torch_dtype,
trust_remote_code=True,
).to(device_id)
server_device = device_id
Expand Down
2 changes: 1 addition & 1 deletion quant_eval/cli/search_rotation.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,7 +14,7 @@

python -m quant_eval.cli.search_rotation \\
--model_name unsloth/Llama-3.2-1B \\
--base_config quant_eval/configs/llama_mxint4.toml \\
--base_config quant_eval/configs/llama_mxint8.toml \\
--calib_data wikitext2 \\
--calib_nsamples 128 \\
--output_json checkpoints/rotation_decisions.json
Expand Down
Original file line number Diff line number Diff line change
@@ -1,4 +1,8 @@
# LLaMA MXInt4: attention + MLP projections
# LLaMA MXInt8 (W8A8): attention + MLP projections
#
# Every block below uses 8-bit weights and 8-bit activations with a
# block size of 32 (MXInt8). Attention projections also quantise the
# bias to 8 bits.

by = "regex_name"

Expand Down
2 changes: 1 addition & 1 deletion quant_eval/eval/llada/eval_llada.py
Original file line number Diff line number Diff line change
Expand Up @@ -104,7 +104,7 @@ def __init__(
config.flash_attention = True
self.model = LLaDAModelLM.from_pretrained(
model_path, trust_remote_code=True,
torch_dtype=torch.bfloat16, config=config, **model_kwargs,
dtype=torch.bfloat16, config=config, **model_kwargs,
)
self.model.eval()

Expand Down
2 changes: 1 addition & 1 deletion quant_eval/scripts/run_lm_eval_phase.sh
Original file line number Diff line number Diff line change
Expand Up @@ -12,7 +12,7 @@ set -euo pipefail

MODEL_NAME="${MODEL_NAME:-Qwen/Qwen2.5-1.5B}"
DEVICE="${DEVICE:-cuda:0}"
QUANT_CONFIG="${QUANT_CONFIG:-quant_eval/configs/llama_mxint4.toml}"
QUANT_CONFIG="${QUANT_CONFIG:-quant_eval/configs/llama_mxint8.toml}"
TASKS="${TASKS:-gsm8k}"
PREFILL_WIDTH="${PREFILL_WIDTH:-16}"
DECODE_WIDTH="${DECODE_WIDTH:-16}"
Expand Down
2 changes: 1 addition & 1 deletion quant_eval/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -118,7 +118,7 @@ def setup_model(model_name, model_parallel, dtype, device, attn_implementation="

model = AutoModelForCausalLM.from_pretrained(
model_name,
torch_dtype=dtype,
dtype=dtype,
attn_implementation=attn_implementation,
trust_remote_code=True,
)
Expand Down