diff --git a/.gitignore b/.gitignore index ea63c21..a9b12fb 100644 --- a/.gitignore +++ b/.gitignore @@ -13,4 +13,7 @@ site/ .cache/ **mase/ **fast-hadamard-transform/ -**calib/ \ No newline at end of file +**calib/ +.venv/ +.pytest_cache/ +*.vcd diff --git a/quant_eval/cli/QUANT_HF_SERVE.md b/quant_eval/cli/QUANT_HF_SERVE.md index 31a12fe..bb6d70d 100644 --- a/quant_eval/cli/QUANT_HF_SERVE.md +++ b/quant_eval/cli/QUANT_HF_SERVE.md @@ -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 diff --git a/quant_eval/cli/eval_dllm.py b/quant_eval/cli/eval_dllm.py index 53e2ad6..9e44cb7 100644 --- a/quant_eval/cli/eval_dllm.py +++ b/quant_eval/cli/eval_dllm.py @@ -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 """ diff --git a/quant_eval/cli/eval_evalplus.py b/quant_eval/cli/eval_evalplus.py index 1418759..78eaee7 100644 --- a/quant_eval/cli/eval_evalplus.py +++ b/quant_eval/cli/eval_evalplus.py @@ -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 @@ -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, diff --git a/quant_eval/cli/eval_llada.py b/quant_eval/cli/eval_llada.py index 9381e05..bd9cce1 100644 --- a/quant_eval/cli/eval_llada.py +++ b/quant_eval/cli/eval_llada.py @@ -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") diff --git a/quant_eval/cli/eval_lm.py b/quant_eval/cli/eval_lm.py index 94fcf7a..225da76 100644 --- a/quant_eval/cli/eval_lm.py +++ b/quant_eval/cli/eval_lm.py @@ -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 """ @@ -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, diff --git a/quant_eval/cli/eval_osworld.py b/quant_eval/cli/eval_osworld.py index e59bd1c..a715c02 100644 --- a/quant_eval/cli/eval_osworld.py +++ b/quant_eval/cli/eval_osworld.py @@ -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 """ diff --git a/quant_eval/cli/eval_phase_bfcl.py b/quant_eval/cli/eval_phase_bfcl.py index d10957a..5283156 100755 --- a/quant_eval/cli/eval_phase_bfcl.py +++ b/quant_eval/cli/eval_phase_bfcl.py @@ -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 \\ @@ -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, diff --git a/quant_eval/cli/eval_phase_lm.py b/quant_eval/cli/eval_phase_lm.py index 5158a15..a5ea1bd 100644 --- a/quant_eval/cli/eval_phase_lm.py +++ b/quant_eval/cli/eval_phase_lm.py @@ -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 @@ -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, diff --git a/quant_eval/cli/eval_ppl.py b/quant_eval/cli/eval_ppl.py index 0310ae3..20423b0 100644 --- a/quant_eval/cli/eval_ppl.py +++ b/quant_eval/cli/eval_ppl.py @@ -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 diff --git a/quant_eval/cli/quant_hf_serve.py b/quant_eval/cli/quant_hf_serve.py index b3c44af..a4cddea 100644 --- a/quant_eval/cli/quant_hf_serve.py +++ b/quant_eval/cli/quant_hf_serve.py @@ -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 @@ -356,7 +356,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, ) @@ -364,7 +364,7 @@ def main( else: model = AutoModelForCausalLM.from_pretrained( model_name, - torch_dtype=torch_dtype, + dtype=torch_dtype, trust_remote_code=True, ).to(device_id) model.eval() @@ -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, ) @@ -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 diff --git a/quant_eval/cli/search_rotation.py b/quant_eval/cli/search_rotation.py index 85f6449..022e87e 100644 --- a/quant_eval/cli/search_rotation.py +++ b/quant_eval/cli/search_rotation.py @@ -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 diff --git a/quant_eval/configs/llama_mxint4.toml b/quant_eval/configs/llama_mxint8.toml similarity index 66% rename from quant_eval/configs/llama_mxint4.toml rename to quant_eval/configs/llama_mxint8.toml index 0f91756..03026d0 100644 --- a/quant_eval/configs/llama_mxint4.toml +++ b/quant_eval/configs/llama_mxint8.toml @@ -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" diff --git a/quant_eval/eval/llada/eval_llada.py b/quant_eval/eval/llada/eval_llada.py index d30d80f..f60b028 100644 --- a/quant_eval/eval/llada/eval_llada.py +++ b/quant_eval/eval/llada/eval_llada.py @@ -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() diff --git a/quant_eval/scripts/run_lm_eval_phase.sh b/quant_eval/scripts/run_lm_eval_phase.sh index 90f9470..a3b2ec3 100755 --- a/quant_eval/scripts/run_lm_eval_phase.sh +++ b/quant_eval/scripts/run_lm_eval_phase.sh @@ -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}" diff --git a/quant_eval/utils.py b/quant_eval/utils.py index 6682e47..c9cd525 100644 --- a/quant_eval/utils.py +++ b/quant_eval/utils.py @@ -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, )