mirror of
https://github.com/NVIDIA/Model-Optimizer.git
synced 2026-10-02 03:14:52 +08:00
[NNBUG: 5701866] Update DS V3.2 PTQ code (#630)
## What does this PR do? **Type of change:** ? Bug fix **Overview:** 1) Update the DS V3.2 repo code reference to the latest version 2) The new DS V3.2 model now includes fp32 layers. We cast it down to match the checkpoint format during loading 3) Fix get_quant_config API change. ## Testing Generate the deepseek-ai/DeepSeek-V3.2 checkpoint ## Before your PR is "*Ready for review*" <!-- If you haven't finished some of the above items you can still open `Draft` PR. --> - **Make sure you read and follow [Contributor guidelines](https://github.com/NVIDIA/TensorRT-Model-Optimizer/blob/main/CONTRIBUTING.md)** and your commits are signed. - **Is this change backward compatible?**: Yes/No <!--- If No, explain why. --> - **Did you write any new necessary tests?**: Yes/No - **Did you add or update any necessary documentation?**: Yes/No - **Did you update [Changelog](https://github.com/NVIDIA/TensorRT-Model-Optimizer/blob/main/CHANGELOG.rst)?**: Yes/No <!--- Only for new features, API changes, critical bug fixes or bw breaking changes. --> ## Additional Information <!-- E.g. related issue. --> Signed-off-by: Chenjie Luo <chenjiel@nvidia.com>
This commit is contained in:
@@ -33,7 +33,7 @@ git clone https://github.com/deepseek-ai/DeepSeek-V3.git && cd DeepSeek-V3 && gi
|
||||
huggingface-cli download deepseek-ai/DeepSeek-V3.2-Exp --local-dir $HF_FP8_CKPT
|
||||
|
||||
# clone DeepSeek-V3.2 Github repository for FP8 inference,
|
||||
git clone https://github.com/deepseek-ai/DeepSeek-V3.2-Exp.git && cd DeepSeek-V3.2-Exp && git checkout 3b99a53
|
||||
git clone https://github.com/deepseek-ai/DeepSeek-V3.2-Exp.git && cd DeepSeek-V3.2-Exp && git checkout 87e509a
|
||||
|
||||
# Install requirements
|
||||
pip install git+https://github.com/Dao-AILab/fast-hadamard-transform.git
|
||||
|
||||
@@ -257,7 +257,18 @@ def load_deepseek_model(model_config: str, model_path: str, batch_size: int):
|
||||
# load model
|
||||
checkpoint_path = os.path.join(model_path, f"model{rank}-mp{world_size}.safetensors")
|
||||
print(f"Loading {checkpoint_path}")
|
||||
|
||||
# Temporary fix for fp32 params
|
||||
fp32_params = {}
|
||||
for name, param in model.named_parameters():
|
||||
if param.dtype == torch.float32 and (
|
||||
"head.weight" in name or "attn.indexer.weights_proj.weight" in name
|
||||
):
|
||||
param.data = param.data.to(torch.get_default_dtype())
|
||||
fp32_params[name] = param
|
||||
load_model(model, checkpoint_path)
|
||||
for param in fp32_params.values():
|
||||
param.data = param.data.to(torch.float32)
|
||||
print(f"Loaded {checkpoint_path}")
|
||||
return model
|
||||
|
||||
@@ -347,7 +358,7 @@ def save_amax_and_quant_config(model, output_path: str, enable_fp8_kvcache: bool
|
||||
# counts = module.activated_expert_counts()
|
||||
# f.writelines(f"{name}: {count}\n" for count in counts)
|
||||
|
||||
quant_config = get_quant_config(model.named_modules())
|
||||
quant_config = get_quant_config(model)
|
||||
|
||||
if enable_fp8_kvcache:
|
||||
quant_config["quantization"]["kv_cache_quant_algo"] = KV_CACHE_FP8
|
||||
|
||||
Reference in New Issue
Block a user