[Security] Enable torch.load(weights_only=True) for secure checkpoint loading + trust_remote_code fix (#1181)

### What does this PR do?

- Add secure checkpoint loading support using
`torch.serialization.add_safe_globals([cls])`. This also removes 1
existing pickle usage.
- Remove hard-coded `trust_remote_code=True`
- Replaces https://github.com/NVIDIA/Model-Optimizer/pull/1056 by
@RinZ27

### Testing
<!-- Mention how have you tested your change if applicable. -->

CICD tests ran

### Before your PR is "*Ready for review*"

Make sure you read and follow [Contributor
guidelines](https://github.com/NVIDIA/Model-Optimizer/blob/main/CONTRIBUTING.md)
and your commits are signed (`git commit -s -S`).

Make sure you read and follow the [Security Best
Practices](https://github.com/NVIDIA/Model-Optimizer/blob/main/SECURITY.md#security-coding-practices-for-contributors)
(e.g. avoiding hardcoded `trust_remote_code=True`, `torch.load(...,
weights_only=False)`, `pickle`, etc.).

- Is this change backward compatible?: ✅ <!--- If ❌, explain why. -->
- If you copied code from any other sources or added a new PIP
dependency, did you follow guidance in `CONTRIBUTING.md`: ✅ <!---
Mandatory -->
- Did you write any new necessary tests?: ✅ <!--- Mandatory for new
features or examples. -->
- Did you update
[Changelog](https://github.com/NVIDIA/Model-Optimizer/blob/main/CHANGELOG.rst)?:
✅ <!--- Only for new features, API changes, critical bug fixes or
backward incompatible changes. -->

### Additional Information

NVBug: 5999336

<!-- This is an auto-generated comment: release notes by coderabbit.ai
-->
## Summary by CodeRabbit

* **New Features**
* Added safe checkpoint save/load helpers and a --trust_remote_code CLI
flag in examples to control remote-code loading.

* **Bug Fixes**
* Checkpoint loading now defaults to safer, weights-only semantics to
reduce arbitrary-code exposure.

* **Documentation**
* CHANGELOG updated with security guidance and opt-in procedure for
unsafe checkpoint loading.

* **Tests**
  * New unit tests validating the safe-load behavior.
<!-- end of auto-generated comment: release notes by coderabbit.ai -->

---------

Signed-off-by: RinZ27 <222222878+RinZ27@users.noreply.github.com>
Signed-off-by: Keval Morabia <28916987+kevalmorabia97@users.noreply.github.com>
Co-authored-by: RinZ27 <222222878+RinZ27@users.noreply.github.com>
This commit is contained in:
Keval Morabia
2026-04-08 00:36:28 +05:30
committed by GitHub
co-authored by RinZ27
parent 80d2f02a2d
commit 5dc17dfd15
39 changed files with 381 additions and 111 deletions
+4
View File
@@ -22,6 +22,10 @@ NVIDIA Model Optimizer Changelog
- Fix Minitron pruning (``mcore_minitron``) for MoE models. Importance estimation hooks were incorrectly registered for MoE modules and NAS step was hanging before this.
**Misc**
- [Security] Changed the default of ``weights_only`` to ``True`` in ``torch.load`` for secure checkpoint loading. If you need to load a checkpoint that requires unpickling arbitrary objects, first register the class in ``torch.serialization.add_safe_globals([cls])`` before loading. Added :meth:`safe_save <modelopt.torch.utils.serialization.safe_save>` and :meth:`safe_load <modelopt.torch.utils.serialization.safe_load>` API to save and load checkpoints securely.
0.43 (2026-04-09)
^^^^^^^^^^^^^^^^^
@@ -95,21 +95,22 @@ def convert_and_save(model, tokenizer, output_path: str):
def create_parser():
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--model_path", type=str, help="path to the fake-quantized model from QAT.")
parser.add_argument(
"--trust_remote_code",
action="store_true",
help="Set trust_remote_code for Huggingface models and tokenizers",
)
parser.add_argument(
"--lora_path",
type=str,
help="path to the LoRA-QAT adapter weights. You can only specify lora_path or model_path, not both.",
)
parser.add_argument(
"--base_path",
type=str,
help="path to the base model used for LoRA-QAT. Only used if lora_path is specified.",
)
parser.add_argument(
"--output_path", type=str, required=True, help="location to save converted model."
)
@@ -121,7 +122,11 @@ if __name__ == "__main__":
parser = create_parser()
args = parser.parse_args()
kwargs = {"device_map": "auto", "torch_dtype": "auto", "trust_remote_code": True}
kwargs = {
"device_map": "auto",
"torch_dtype": "auto",
"trust_remote_code": args.trust_remote_code,
}
if args.lora_path:
assert args.model_path is None, "You can only specify lora_path or model_path, not both."
model_path = args.base_path
@@ -140,7 +145,7 @@ if __name__ == "__main__":
gc.collect()
# Load tokenizer
tokenizer = AutoTokenizer.from_pretrained(model_path, trust_remote_code=True)
tokenizer = AutoTokenizer.from_pretrained(model_path, trust_remote_code=args.trust_remote_code)
# Quantize and save model
convert_and_save(model, tokenizer, args.output_path)
+9 -2
View File
@@ -128,10 +128,11 @@ def modelopt_ptq(
auto_quantize_bits: float | None = None,
calib_dataset: str = "cnn_dailymail",
calib_batch_size: int = 8,
trust_remote_code: bool = False,
) -> torch.nn.Module:
"""Quantize the model with modelopt."""
model = AutoModelForCausalLM.from_pretrained(
model_path, trust_remote_code=True, torch_dtype="auto", device_map="auto"
model_path, trust_remote_code=trust_remote_code, torch_dtype="auto", device_map="auto"
)
model.eval()
@@ -139,7 +140,7 @@ def modelopt_ptq(
model_path,
model_max_length=2048,
padding_side="left",
trust_remote_code=True,
trust_remote_code=trust_remote_code,
)
# sanitize tokenizer
if tokenizer.pad_token != "<unk>":
@@ -213,6 +214,11 @@ if __name__ == "__main__":
"regular quantization without auto_quantize search will be applied."
),
)
parser.add_argument(
"--trust_remote_code",
action="store_true",
help="Set trust_remote_code for Huggingface models and tokenizers",
)
args = parser.parse_args()
@@ -223,4 +229,5 @@ if __name__ == "__main__":
args.num_samples,
auto_quantize_bits=args.effective_bits,
calib_batch_size=args.calib_batch_size,
trust_remote_code=args.trust_remote_code,
)
+1 -2
View File
@@ -38,6 +38,7 @@
# limitations under the License.
import warnings
import datasets
from lm_eval import utils
from lm_eval.__main__ import cli_evaluate, parse_eval_args, setup_parser
from lm_eval.api.model import T
@@ -180,8 +181,6 @@ if __name__ == "__main__":
model_args = utils.simple_parse_args_string(args.model_args)
if args.trust_remote_code:
import datasets
datasets.config.HF_DATASETS_TRUST_REMOTE_CODE = True
model_args["trust_remote_code"] = True
args.trust_remote_code = None
+15 -8
View File
@@ -74,6 +74,7 @@ from transformers import (
class EvalModel(BaseModel, arbitrary_types_allowed=True):
model_path: str
trust_remote_code: bool = False
max_input_length: int = 512
max_output_length: int = 512
dtype: str = "auto"
@@ -92,7 +93,6 @@ class EvalModel(BaseModel, arbitrary_types_allowed=True):
class OpenAIModel(EvalModel):
model_path: str
engine: str = ""
use_azure: bool = False
tokenizer: tiktoken.Encoding | None
@@ -173,7 +173,6 @@ class OpenAIModel(EvalModel):
class SeqToSeqModel(EvalModel):
model_path: str
model: PreTrainedModel | None = None
tokenizer: PreTrainedTokenizer | None = None
lora_path: str = ""
@@ -191,7 +190,9 @@ class SeqToSeqModel(EvalModel):
args.update(torch_dtype=getattr(torch, self.dtype) if self.dtype != "auto" else "auto")
if self.attn_implementation:
args["attn_implementation"] = self.attn_implementation
self.model = AutoModelForSeq2SeqLM.from_pretrained(self.model_path, **args)
self.model = AutoModelForSeq2SeqLM.from_pretrained(
self.model_path, trust_remote_code=self.trust_remote_code, **args
)
print_gpu_utilization()
if self.lora_path:
self.model = PeftModel.from_pretrained(self.model, self.lora_path)
@@ -199,7 +200,9 @@ class SeqToSeqModel(EvalModel):
if "device_map" not in args:
self.model.to(self.device)
if self.tokenizer is None:
self.tokenizer = AutoTokenizer.from_pretrained(self.model_path)
self.tokenizer = AutoTokenizer.from_pretrained(
self.model_path, trust_remote_code=self.trust_remote_code
)
def run(self, prompt: str, **kwargs) -> str:
self.load()
@@ -247,7 +250,7 @@ class CausalModel(SeqToSeqModel):
if self.attn_implementation:
args["attn_implementation"] = self.attn_implementation
self.model = AutoModelForCausalLM.from_pretrained(
self.model_path, trust_remote_code=True, **args
self.model_path, trust_remote_code=self.trust_remote_code, **args
)
self.model.eval()
if "device_map" not in args:
@@ -256,7 +259,9 @@ class CausalModel(SeqToSeqModel):
# Sampling with temperature will cause MMLU to drop
self.model.generation_config.do_sample = False
if self.tokenizer is None:
self.tokenizer = AutoTokenizer.from_pretrained(self.model_path, trust_remote_code=True)
self.tokenizer = AutoTokenizer.from_pretrained(
self.model_path, trust_remote_code=self.trust_remote_code
)
def run(self, prompt: str, **kwargs) -> str:
self.load()
@@ -487,10 +492,12 @@ class GPTQModel(LlamaModel):
class ChatGLMModel(SeqToSeqModel):
def load(self):
if self.tokenizer is None:
self.tokenizer = AutoTokenizer.from_pretrained(self.model_path, trust_remote_code=True)
self.tokenizer = AutoTokenizer.from_pretrained(
self.model_path, trust_remote_code=self.trust_remote_code
)
if self.model is None:
self.model = AutoModel.from_pretrained(
self.model_path, trust_remote_code=True
self.model_path, trust_remote_code=self.trust_remote_code
).half() # FP16 is required for ChatGLM
self.model.eval()
self.model.to(self.device)
+10 -4
View File
@@ -53,7 +53,13 @@ SPECULATIVE_MODEL_LIST = ["Eagle", "Medusa"]
def run_nemotron_vl_preview(
full_model, tokenizer, input_ids, pyt_ckpt_path, stage_name, allow_fallback=False
full_model,
tokenizer,
input_ids,
pyt_ckpt_path,
stage_name,
allow_fallback=False,
trust_remote_code=False,
):
"""Run text-only and VL preview generation for Nemotron VL models.
@@ -64,7 +70,7 @@ def run_nemotron_vl_preview(
pyt_ckpt_path: Path to the model checkpoint
stage_name: Description of the stage (e.g., "before quantization", "after quantization")
allow_fallback: Whether to allow fallback to standard generate on failure
trust_remote_code: Whether to trust remote code for Huggingface models and tokenizers
Returns:
Generated text response or None if generation failed
"""
@@ -80,7 +86,7 @@ def run_nemotron_vl_preview(
# Try text-only generation (may fail for encoder-decoder models like Nemotron-Parse)
text_response = run_text_only_generation(
full_model, tokenizer, question, generation_config, pyt_ckpt_path
full_model, tokenizer, question, generation_config, pyt_ckpt_path, trust_remote_code
)
generated_ids = None
@@ -93,7 +99,7 @@ def run_nemotron_vl_preview(
# Run additional VL test with images
print(f"Running additional VL test with images ({stage_name})...")
run_vl_preview_generation(full_model, tokenizer, pyt_ckpt_path, stage_name)
run_vl_preview_generation(full_model, tokenizer, pyt_ckpt_path, stage_name, trust_remote_code)
return generated_ids
+2
View File
@@ -770,6 +770,7 @@ def pre_quantize(
args.pyt_ckpt_path,
"before quantization",
allow_fallback=False,
trust_remote_code=args.trust_remote_code,
)
else:
generated_ids_before_ptq = full_model.generate(preview_input_ids, max_new_tokens=100)
@@ -820,6 +821,7 @@ def post_quantize(
args.pyt_ckpt_path,
"after quantization",
allow_fallback=False,
trust_remote_code=args.trust_remote_code,
)
else:
warnings.warn(
+15 -7
View File
@@ -21,7 +21,7 @@ from PIL import Image
from transformers import AutoImageProcessor, AutoProcessor
def run_vl_preview_generation(model, tokenizer, model_path, stage_name):
def run_vl_preview_generation(model, tokenizer, model_path, stage_name, trust_remote_code=False):
"""Run preview generation for VL models using sample images.
Args:
@@ -29,7 +29,7 @@ def run_vl_preview_generation(model, tokenizer, model_path, stage_name):
tokenizer: The tokenizer
model_path: Path to the model (for loading image processor)
stage_name: Description of the stage (e.g., "before quantization")
trust_remote_code: Whether to trust remote code for Huggingface models and tokenizers
Returns:
Generated response text for logging/comparison
"""
@@ -85,7 +85,9 @@ def run_vl_preview_generation(model, tokenizer, model_path, stage_name):
# Try to detect the VL model has chat method or generate method
if hasattr(model, "chat"):
image_processor = AutoImageProcessor.from_pretrained(model_path, trust_remote_code=True)
image_processor = AutoImageProcessor.from_pretrained(
model_path, trust_remote_code=trust_remote_code
)
image_features = image_processor([image]) # Pass as list with single image
@@ -103,7 +105,9 @@ def run_vl_preview_generation(model, tokenizer, model_path, stage_name):
**image_features,
)
else:
processor = AutoProcessor.from_pretrained(model_path, trust_remote_code=True)
processor = AutoProcessor.from_pretrained(
model_path, trust_remote_code=trust_remote_code
)
# Use chat template if available, otherwise fall back to default task prompt
if hasattr(tokenizer, "chat_template") and tokenizer.chat_template is not None:
@@ -190,7 +194,9 @@ def run_vl_preview_generation(model, tokenizer, model_path, stage_name):
return None
def run_text_only_generation(model, tokenizer, question, generation_config, model_path):
def run_text_only_generation(
model, tokenizer, question, generation_config, model_path, trust_remote_code=False
):
"""Run text-only generation for VL models, supporting both chat and generate methods.
Args:
@@ -199,7 +205,7 @@ def run_text_only_generation(model, tokenizer, question, generation_config, mode
question: The text question to ask
generation_config: Generation configuration
model_path: Path to the model (for loading processor if needed)
trust_remote_code: Whether to trust remote code for Huggingface models and tokenizers
Returns:
Generated response text or None if failed
"""
@@ -209,7 +215,9 @@ def run_text_only_generation(model, tokenizer, question, generation_config, mode
response = model.chat(tokenizer, None, question, generation_config, history=None)
return response
else:
processor = AutoProcessor.from_pretrained(model_path, trust_remote_code=True)
processor = AutoProcessor.from_pretrained(
model_path, trust_remote_code=trust_remote_code
)
# Create text-only messages
messages = [
@@ -30,14 +30,14 @@ TRAIN_RATIO, VALID_RATIO = 0.95, 0.025
_TOKENIZER = None
def init_tokenizer(name: str) -> None:
def init_tokenizer(name: str, trust_remote_code: bool = False) -> None:
"""Load HuggingFace tokenizer for chat template."""
global _TOKENIZER
if name:
from transformers import AutoTokenizer
print(f"Loading tokenizer: {name}")
_TOKENIZER = AutoTokenizer.from_pretrained(name, trust_remote_code=True)
_TOKENIZER = AutoTokenizer.from_pretrained(name, trust_remote_code=trust_remote_code)
def format_text(messages: list[dict], reasoning: str = "") -> str:
@@ -159,10 +159,15 @@ def main():
p.add_argument(
"--include-reasoning", action="store_true", help="Include COT for Thinking models"
)
p.add_argument(
"--trust_remote_code",
action="store_true",
help="Set trust_remote_code for Huggingface models and tokenizers",
)
args = p.parse_args()
if args.tokenizer:
init_tokenizer(args.tokenizer)
init_tokenizer(args.tokenizer, args.trust_remote_code)
# Build suffix
suffix = f"{int(args.sample_percent)}pct"
@@ -85,6 +85,11 @@ def parse_args() -> argparse.Namespace:
default=1,
help="""Data parallel world size. Number of tasks on SLURM.""",
)
parser.add_argument(
"--trust_remote_code",
action="store_true",
help="Set trust_remote_code for Huggingface models and tokenizers",
)
return parser.parse_args()
@@ -130,11 +135,11 @@ def main(args: argparse.Namespace) -> None:
dataset = dataset.select(range(args.debug_max_num_conversations))
model = AutoModel.from_pretrained(
args.model, torch_dtype="auto", device_map="auto", trust_remote_code=True
args.model, torch_dtype="auto", device_map="auto", trust_remote_code=args.trust_remote_code
)
num_hidden_layers = getattr(model.config, "num_hidden_layers", None)
tokenizer = AutoTokenizer.from_pretrained(args.model, trust_remote_code=True)
tokenizer = AutoTokenizer.from_pretrained(args.model, trust_remote_code=args.trust_remote_code)
if tokenizer.pad_token is None:
tokenizer.pad_token = tokenizer.eos_token
tokenizer.chat_template = tokenizer.chat_template.replace(REMOVE_THINK_CHAT_TEMPLATE, "")
@@ -55,6 +55,7 @@ def validate_ar(model, tokenizer, ds, steps=3, osl=20, num_samples=80, device=No
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--model_path", type=str, required=True, help="Path to model directory")
parser.add_argument("--trust_remote_code", action="store_true", help="Trust remote code")
parser.add_argument("--steps", type=int, default=3, help="Steps for AR validation")
parser.add_argument(
"--osl", type=int, default=32, help="Output sequence length for AR validation"
@@ -72,8 +73,12 @@ def main():
accelerator = Accelerator()
# Load model and tokenizer
model = load_vlm_or_llm(args.model_path, device_map="auto")
tokenizer = AutoTokenizer.from_pretrained(args.model_path)
model = load_vlm_or_llm(
args.model_path, device_map="auto", trust_remote_code=args.trust_remote_code
)
tokenizer = AutoTokenizer.from_pretrained(
args.model_path, trust_remote_code=args.trust_remote_code
)
model.eval()
model = accelerator.prepare(model)
@@ -29,6 +29,7 @@ def parse_args():
description="Export a HF checkpoint (with ModelOpt state) for deployment."
)
parser.add_argument("--model_path", type=str, default="Path of the trained checkpoint.")
parser.add_argument("--trust_remote_code", action="store_true", help="Trust remote code")
parser.add_argument(
"--export_path", type=str, default="Destination directory for exported files."
)
@@ -38,11 +39,10 @@ def parse_args():
mto.enable_huggingface_checkpointing()
args = parse_args()
model = load_vlm_or_llm(args.model_path, torch_dtype="auto")
model = load_vlm_or_llm(
args.model_path, torch_dtype="auto", trust_remote_code=args.trust_remote_code
)
model.eval()
with torch.inference_mode():
export_speculative_decoding(
model,
export_dir=args.export_path,
)
export_speculative_decoding(model, export_dir=args.export_path)
print(f"Exported checkpoint to {args.export_path}")
@@ -55,6 +55,11 @@ def parse_args() -> argparse.Namespace:
"the local serving engine. This should match the value used by the server."
),
)
parser.add_argument(
"--trust_remote_code",
action="store_true",
help="Set trust_remote_code for Huggingface models and tokenizers",
)
## Client Parameters ##
parser.add_argument(
"--base-url",
@@ -133,7 +138,9 @@ async def main(args: argparse.Namespace) -> None:
base_url=args.base_url,
)
tokenizer = AutoTokenizer.from_pretrained(args.model_card, trust_remote_code=True)
tokenizer = AutoTokenizer.from_pretrained(
args.model_card, trust_remote_code=args.trust_remote_code
)
if tokenizer.pad_token is None:
tokenizer.pad_token = tokenizer.eos_token
bos_token_id = tokenizer.bos_token_id
@@ -501,7 +501,11 @@ def main(
tokenizer = get_tokenizer(model_ckpt_path, trust_remote_code=trust_remote_code)
model = select_model(
max_input_length=MAX_SEQ_LEN, max_output_length=2, dtype=dtype, **kwargs
max_input_length=MAX_SEQ_LEN,
max_output_length=2,
dtype=dtype,
trust_remote_code=trust_remote_code,
**kwargs,
)
assert isinstance(model, EvalModel)
if quant_cfg:
@@ -49,6 +49,7 @@ class EvalModel(BaseModel):
model_config = ConfigDict(arbitrary_types_allowed=True)
model_path: str
trust_remote_code: bool = False
max_input_length: int = 512
max_output_length: int = 512
dtype: str = "auto"
@@ -84,7 +85,9 @@ class SeqToSeqModel(EvalModel):
args.update(torch_dtype=getattr(torch, self.dtype))
else:
args.update(torch_dtype="auto")
self.model = AutoModelForSeq2SeqLM.from_pretrained(self.model_path, **args)
self.model = AutoModelForSeq2SeqLM.from_pretrained(
self.model_path, trust_remote_code=self.trust_remote_code, **args
)
print_gpu_utilization()
if self.lora_path:
self.model = PeftModel.from_pretrained(self.model, self.lora_path)
@@ -92,7 +95,9 @@ class SeqToSeqModel(EvalModel):
if "device_map" not in args:
self.model.to(self.device)
if self.tokenizer is None:
self.tokenizer = AutoTokenizer.from_pretrained(self.model_path)
self.tokenizer = AutoTokenizer.from_pretrained(
self.model_path, trust_remote_code=self.trust_remote_code
)
def run(self, prompt: str, **kwargs) -> str:
self.load()
@@ -143,7 +148,7 @@ class CausalModel(SeqToSeqModel):
args.update(device_map="auto", load_in_8bit=True)
args.update(torch_dtype=getattr(torch, self.dtype) if self.dtype != "auto" else "auto")
self.model = AutoModelForCausalLM.from_pretrained(
self.model_path, trust_remote_code=True, **args
self.model_path, trust_remote_code=self.trust_remote_code, **args
)
self.model.eval()
if "device_map" not in args:
@@ -152,7 +157,9 @@ class CausalModel(SeqToSeqModel):
# Sampling with temperature will cause MMLU to drop
self.model.generation_config.do_sample = False
if self.tokenizer is None:
self.tokenizer = AutoTokenizer.from_pretrained(self.model_path, trust_remote_code=True)
self.tokenizer = AutoTokenizer.from_pretrained(
self.model_path, trust_remote_code=self.trust_remote_code
)
def run(self, prompt: str, **kwargs) -> str:
self.load()
@@ -200,7 +207,7 @@ class AutoAWQCausalModel(SeqToSeqModel):
args.update(device_map="auto", load_in_8bit=True)
args.update(torch_dtype=getattr(torch, self.dtype) if self.dtype != "auto" else "auto")
self.model = AutoAWQForCausalLM.from_quantized(
self.model_path, trust_remote_code=True, **args
self.model_path, trust_remote_code=self.trust_remote_code, **args
)
self.model.eval()
if "device_map" not in args:
@@ -209,7 +216,9 @@ class AutoAWQCausalModel(SeqToSeqModel):
# Sampling with temperature will cause MMLU to drop
self.model.config.do_sample = False
if self.tokenizer is None:
self.tokenizer = AutoTokenizer.from_pretrained(self.model_path, trust_remote_code=True)
self.tokenizer = AutoTokenizer.from_pretrained(
self.model_path, trust_remote_code=self.trust_remote_code
)
def run(self, prompt: str, **kwargs) -> str:
self.load()
+1 -1
View File
@@ -174,7 +174,7 @@ These scripts are currently validated with following settings:
- Calibration size - 32
- Calibration EPs - \[`cuda`, `cpu`\]
- Audio dataset - `librispeech_asr` dataset (32 samples used for calibration, 100+ samples used for WER test)
- `load_dataset("librispeech_asr", "clean", split="test", trust_remote_code=True)`
- `load_dataset("librispeech_asr", "clean", split="test")`
- Quantization support for various ONNX files - `encoder_model.onnx`, `decoder_model.onnx`, `decoder_with_past_model.onnx`
- The `use_merged` argument in optimum-ORT's Whisper model API is kept False.
@@ -275,7 +275,7 @@ def main(args):
processor = WhisperProcessor.from_pretrained(args.model_name, cache_dir=args.cache_dir)
asr_dataset = load_dataset("librispeech_asr", "clean", split="test", trust_remote_code=True)
asr_dataset = load_dataset("librispeech_asr", "clean", split="test")
# asr_dataset = load_dataset("librispeech_asr", "all", split="test.clean")
calib_data = None
@@ -85,9 +85,7 @@ def main(args):
print(f"\n\n-- Content of input audio-file = {prediction}\n\n")
if args.run_wer_test:
librispeech_test_clean = load_dataset(
"librispeech_asr", "clean", split="test", trust_remote_code=True
)
librispeech_test_clean = load_dataset("librispeech_asr", "clean", split="test")
references = []
predictions = []
@@ -65,6 +65,7 @@ import modelopt.torch.opt as mto
import modelopt.torch.quantization as mtq
from modelopt.torch.distill.distillation_model import DistillationModel
from modelopt.torch.quantization.config import NVFP4_DEFAULT_CFG
from modelopt.torch.utils import safe_load
logger = logging.getLogger(__name__)
@@ -131,13 +132,13 @@ def detect_format(path: str) -> str:
return "safetensors"
def load_state_dict_any_format(path: str, label: str = "") -> tuple[dict, dict | None]:
def load_state_dict_any_format(path: str, label: str = "", **kwargs) -> tuple[dict, dict | None]:
"""Load state dict from either torch pickle or safetensors."""
fmt = detect_format(path)
logger.info(f"[{label}] Detected format: {fmt} for {path}")
if fmt == "torch":
raw = torch.load(path, map_location="cpu", weights_only=False)
raw = safe_load(path, map_location="cpu", **kwargs)
if isinstance(raw, dict) and "state_dict" in raw:
return raw["state_dict"], None
return raw, None
+3 -3
View File
@@ -25,6 +25,7 @@ from typing import Any
import torch
from modelopt.torch.utils import distributed as dist
from modelopt.torch.utils import safe_load
from .model_config_utils import (
model_config_from_dict,
@@ -41,7 +42,7 @@ class NFSWorkspace:
communication nor barrier. It is users' responsibility to synchronize
all ranks (local and remove processes).
This implementation uses `torch.save` and `torch.load` for serialization.
This implementation uses `torch.save` and `safe_load` (`torch.load(weights_only=True)`) for serialization.
Args:
workspace_path: the path to the NFS directory for postprocess cross rank communication.
@@ -91,8 +92,7 @@ class NFSWorkspace:
raise ValueError("NFSWorkspace is not initialized!")
state_path = self._get_state_path(target_rank)
if state_path.exists():
# Security NOTE: weights_only=False is used here on ModelOpt-generated ckpt, not on untrusted user input
state = torch.load(state_path, map_location="cpu", weights_only=False)
state = safe_load(state_path, map_location="cpu")
return state["config"], state["weight"]
else:
return None, None
+23
View File
@@ -633,3 +633,26 @@ class ModelConfig:
def hidden_act(self):
"""Returns the hidden_act of the model."""
return self.layers[0].mlp.hidden_act
# Register all config classes as safe globals
torch.serialization.add_safe_globals(
[
EmbeddingConfig,
LayernormConfig,
LinearConfig,
LinearActConfig,
ConvConfig,
QKVConfig,
RelativeAttentionTableConfig,
AttentionConfig,
MLPConfig,
ExpertConfig,
RgLruConfig,
RecurrentConfig,
MOEConfig,
DecoderLayerConfig,
MedusaHeadConfig,
ModelConfig,
]
)
+9
View File
@@ -20,6 +20,7 @@ import json
from collections.abc import Callable, ItemsView, Iterator, KeysView, ValuesView
from typing import Any, TypeAlias
import torch
from pydantic import (
BaseModel,
Field,
@@ -65,6 +66,14 @@ class ModeloptBaseConfig(BaseModel):
model_config = PyDanticConfigDict(extra="forbid", validate_assignment=True)
def __init_subclass__(cls, **kwargs: Any) -> None:
"""Register the config class as a safe global for torch serialization.
It can be used to load the config with torch.load(weights_only=True) in safe_load().
"""
super().__init_subclass__(**kwargs)
torch.serialization.add_safe_globals([cls])
def model_dump(self, **kwargs):
"""Dump the config to a dictionary with aliases and no warnings by default."""
kwargs = {"by_alias": True, "warnings": False, **kwargs}
+5 -9
View File
@@ -34,7 +34,7 @@ import torch
import torch.nn as nn
from modelopt import __version__
from modelopt.torch.utils import ModelLike, init_model_from_model_like, unwrap_model
from modelopt.torch.utils import ModelLike, init_model_from_model_like, safe_load, unwrap_model
from .config import ConfigDict, ModeloptBaseConfig
from .mode import (
@@ -523,12 +523,8 @@ def load_modelopt_state(modelopt_state_path: str | os.PathLike, **kwargs) -> dic
Returns:
A modelopt state dictionary describing the modifications to the model.
"""
# Security NOTE: weights_only=False is used here on ModelOpt-generated state_dict, not on untrusted user input
kwargs.setdefault("weights_only", False)
kwargs.setdefault("map_location", "cpu")
# TODO: Add some validation to ensure the file is a valid modelopt state file.
modelopt_state = torch.load(modelopt_state_path, **kwargs)
return modelopt_state
return safe_load(modelopt_state_path, **kwargs)
def restore_from_modelopt_state(
@@ -545,12 +541,13 @@ def restore_from_modelopt_state(
.. code-block:: python
import modelopt.torch.opt as mto
from modelopt.torch.utils import safe_load
model = ... # Create the model-like object
# Restore the previously saved modelopt state followed by model weights
mto.restore_from_modelopt_state(model, modelopt_state_path="modelopt_state.pt")
model.load_state_dict(torch.load("model_weights.pt"), ...) # Load the model weights
model.load_state_dict(safe_load("model_weights.pt"), ...) # Load the model weights
If you want to restore the model weights and the modelopt state with saved scales, please use
:meth:`mto.restore()<modelopt.torch.opt.conversion.restore>`.
@@ -628,8 +625,7 @@ def restore(model: ModelLike, f: str | os.PathLike | BinaryIO, **kwargs) -> nn.M
# load checkpoint
kwargs.setdefault("map_location", "cpu")
kwargs.setdefault("weights_only", False)
objs = torch.load(f, **kwargs)
objs = safe_load(f, **kwargs)
# restore model architecture
model_restored = restore_from_modelopt_state(model, objs["modelopt_state"])
+8
View File
@@ -28,6 +28,14 @@ __all__ = ["HPType", "Hparam"]
class CustomHPType(ABC):
"""Custom hyperparameter type base class for user-defined hparam types."""
def __init_subclass__(cls, **kwargs) -> None:
"""Register the custom hparam type as a safe global for torch serialization.
It can be used to load the custom hparam type with torch.load(weights_only=True) in safe_load().
"""
super().__init_subclass__(**kwargs)
torch.serialization.add_safe_globals([cls])
@abstractmethod
def __repr__(self) -> str:
"""Return string representation with relevant properties of the class."""
@@ -31,6 +31,7 @@ from megatron.core.transformer.module import Float16Module
import modelopt
import modelopt.torch.opt as mto
import modelopt.torch.utils.distributed as dist
from modelopt.torch.utils import safe_load
from modelopt.torch.utils.network import SUPPORTED_WRAPPERS
SUPPORTED_WRAPPERS[Float16Module] = "module"
@@ -203,10 +204,7 @@ def restore_sharded_modelopt_state(
return
# Loading the common modelopt_state (replicated on all ranks)
# Security NOTE: weights_only=False is used here on NVIDIA-generated file, not on untrusted user input
common_modelopt_state = torch.load(
modelopt_checkpoint_name + "/" + COMMON_STATE_FNAME, weights_only=False
)
common_modelopt_state = safe_load(modelopt_checkpoint_name + "/" + COMMON_STATE_FNAME)
modelopt_load_version = common_modelopt_state["modelopt_version"]
+8 -7
View File
@@ -15,8 +15,8 @@
"""Support quantization and save/resore for Megatron."""
import contextlib
import pickle # nosec
import types
from io import BytesIO
from typing import Any
import megatron.core.transformer.mlp as megatron_mlp
@@ -24,6 +24,8 @@ import regex as re
import torch
from megatron.core.parallel_state import get_data_parallel_group
from modelopt.torch.utils import safe_load
from ..dynamic import DynamicModule
@@ -82,8 +84,10 @@ def _modelopt_get_extra_state(self):
# Serialize state into byte tensor
torch.cuda.synchronize()
state_serialized = bytearray(pickle.dumps(extra_state)) # nosec
state_serialized = torch.frombuffer(state_serialized, dtype=torch.uint8)
# Use torch.save for serialization to match safe_load
buffer = BytesIO()
torch.save(extra_state, buffer)
state_serialized = torch.frombuffer(buffer.getvalue(), dtype=torch.uint8)
return state_serialized
@@ -102,10 +106,7 @@ def _modelopt_set_extra_state(self, state: Any):
if state.numel() == 0:
return
# Default format: byte tensor with pickled data
#
# TODO: possible deserialization improvement
# https://github.com/NVIDIA/TensorRT-LLM/blob/main/tensorrt_llm/serialization.py
extra_state = pickle.loads(state.detach().cpu().numpy().tobytes()) # nosec
extra_state = safe_load(state.detach().cpu().numpy().tobytes())
else:
raise RuntimeError("Unsupported extra_state format.")
+3 -4
View File
@@ -21,7 +21,7 @@ from collections.abc import Callable
import torch
from peft import PeftModel
from modelopt.torch.utils import get_unwrapped_name, print_rank_0
from modelopt.torch.utils import get_unwrapped_name, print_rank_0, safe_load
from ..conversion import ModeloptStateManager, modelopt_state, restore_from_modelopt_state
from .huggingface import register_for_patching
@@ -83,9 +83,8 @@ def _new_load_adapter(self, model_id, adapter_name, *args, **kwargs):
if os.path.isfile(_get_quantizer_state_save_path(model_id)):
from modelopt.torch.quantization.nn import TensorQuantizer
# Security NOTE: weights_only=False is used here on ModelOpt-generated state_dict, not on untrusted user input
quantizer_state_dict = torch.load(
_get_quantizer_state_save_path(model_id), map_location="cpu", weights_only=False
quantizer_state_dict = safe_load(
_get_quantizer_state_save_path(model_id), map_location="cpu"
)
for name, module in self.named_modules():
if isinstance(module, TensorQuantizer):
+10 -5
View File
@@ -30,11 +30,17 @@ from typing import TYPE_CHECKING, Any, final
import numpy as np
import pulp
import torch
import torch.nn as nn
from modelopt.torch.utils import distributed as dist
from modelopt.torch.utils import no_stdout, print_rank_0, run_forward_loop, warn_rank_0
from modelopt.torch.utils import (
no_stdout,
print_rank_0,
run_forward_loop,
safe_load,
safe_save,
warn_rank_0,
)
if TYPE_CHECKING:
from pathlib import Path
@@ -276,8 +282,7 @@ class BaseSearcher(ABC):
return False
print_rank_0(f"Loading searcher state from {checkpoint}...")
# Security NOTE: weights_only=False is used here on ModelOpt-generated ckpt, not on untrusted user input
state_dict = torch.load(checkpoint, weights_only=False)
state_dict = safe_load(checkpoint)
if strict:
assert state_dict.keys() == self.state_dict().keys(), "Keys in checkpoint don't match!"
for key, default_val in self.default_state_dict.items():
@@ -301,7 +306,7 @@ class BaseSearcher(ABC):
save_dirname, _ = os.path.split(checkpoint)
if save_dirname:
os.makedirs(save_dirname, exist_ok=True)
torch.save(self.state_dict(), checkpoint)
safe_save(self.state_dict(), checkpoint)
class LPS:
@@ -26,7 +26,7 @@ from omegaconf import DictConfig, OmegaConf
from torch import nn
import modelopt.torch.utils.distributed as dist
from modelopt.torch.utils import json_dump
from modelopt.torch.utils import json_dump, safe_load
__all__ = [
"ForwardHook",
@@ -734,11 +734,7 @@ class LayerNormContributionHook(ForwardHook):
all_scores = []
for activation_file in activation_files:
print(f"Loading activations from {activation_file}")
# SECURITY: weights_only=False is required because files contain dictionaries with tensors.
# These files are generated by dump_activations_logs() in this module and contain
# hook state dictionaries. The activations_log_dir should only contain trusted files
# generated by the same codebase, not from untrusted sources.
activation_data = torch.load(activation_file, map_location="cpu", weights_only=False)
activation_data = safe_load(activation_file, map_location="cpu")
# Extract scores from the activation data
for module_name, hook_data in activation_data.items():
@@ -72,6 +72,8 @@ import torch
import torch.nn as nn
import torch.nn.functional as F
from modelopt.torch.utils import safe_load
class OutputSaveHook:
"""Hook to capture and save module outputs during forward pass."""
@@ -180,21 +182,20 @@ def main():
default=None,
help="Path to save comparison statistics as JSON",
)
parser.add_argument(
"--no-weights-only",
action="store_true",
help="Do not use weights_only=True when loading the data to allow unpickling arbitrary objects",
)
args = parser.parse_args()
# Load reference data
print(f"\nLoading reference: {args.reference}")
# SECURITY: weights_only=False is required because files contain dictionaries with tensors.
# These files are expected to be generated by save_multi_layer_outputs() in this module,
# not from untrusted sources. Users should only load files they generated themselves.
ref_data = torch.load(args.reference, map_location="cpu", weights_only=False)
ref_data = safe_load(args.reference, weights_only=not args.no_weights_only)
# Load comparison data
print(f"Loading compare: {args.compare}")
# SECURITY: weights_only=False is required because files contain dictionaries with tensors.
# These files are expected to be generated by save_multi_layer_outputs() in this module,
# not from untrusted sources. Users should only load files they generated themselves.
comp_data = torch.load(args.compare, map_location="cpu", weights_only=False)
comp_data = safe_load(args.compare, weights_only=not args.no_weights_only)
# Compare multi-layer outputs
compare_multi_layer(ref_data, comp_data, args.output_json)
@@ -15,6 +15,10 @@
"""Abstract base class for calibrators."""
from typing import Any
import torch
__all__ = ["_Calibrator"]
@@ -36,6 +40,14 @@ class _Calibrator:
self._axis = axis
self._unsigned = unsigned
def __init_subclass__(cls, **kwargs: Any) -> None:
"""Register the calibrator classes as a safe global for torch serialization.
It can be used to load the calibrator with torch.load(weights_only=True) in safe_load().
"""
super().__init_subclass__(**kwargs)
torch.serialization.add_safe_globals([cls])
def collect(self, x):
"""Abstract method: collect tensor statistics used to compute amax.
@@ -16,6 +16,7 @@
"""Base Class for Real Quantized Tensor."""
import enum
from typing import Any
import torch
from torch.distributed.fsdp._fully_shard._fsdp_param import FSDPParam
@@ -63,6 +64,14 @@ class BaseQuantizedTensor:
}
self._quantized_data = quantized_data
def __init_subclass__(cls, **kwargs: Any) -> None:
"""Register the quantized tensor class as a safe global for torch serialization.
It can be used to load the quantized tensor with torch.load(weights_only=True) in safe_load().
"""
super().__init_subclass__(**kwargs)
torch.serialization.add_safe_globals([cls])
@classmethod
def quantize(cls, input: torch.Tensor, block_size: int):
"""Pack a fake torch.Tensor into a real quantized tensor.
@@ -227,3 +236,7 @@ def pack_real_quantize_weight(module, force_quantize: bool = False):
if name != "":
with fsdp2_aware_weight_update(module, m):
_compress_and_update_module_weight(m)
# Register QTensorWrapper as safe global
torch.serialization.add_safe_globals([QTensorWrapper])
+1
View File
@@ -26,5 +26,6 @@ from .network import *
from .perf import *
from .regex import *
from .robust_json import *
from .serialization import *
from .tensor import *
from .vlm_dataset_utils import *
+67
View File
@@ -0,0 +1,67 @@
# SPDX-FileCopyrightText: Copyright (c) 2023-2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Serialization utilities for secure checkpoint saving and loading."""
import os
from collections import Counter, OrderedDict
from io import BytesIO
from typing import Any, BinaryIO
import torch
_SAFE_DICT_TYPES = (dict, OrderedDict, Counter)
def _sanitize_for_save(obj: Any) -> Any:
"""Recursively convert container subclasses to types accepted by ``weights_only=True``.
* ``dict`` subclasses not in {dict, OrderedDict, Counter} e.g. defaultdict → plain ``dict``
* ``list`` subclasses → plain ``list``
* Recurses into dict values, list/tuple elements.
* Leaves tensors, scalars, strings, bytes, etc. untouched.
"""
if isinstance(obj, dict):
sanitized = {k: _sanitize_for_save(v) for k, v in obj.items()}
if type(obj) in _SAFE_DICT_TYPES:
return type(obj)(sanitized)
return sanitized
if isinstance(obj, list):
sanitized_list = [_sanitize_for_save(v) for v in obj]
if type(obj) is list:
return sanitized_list
return sanitized_list
if isinstance(obj, tuple):
return tuple(_sanitize_for_save(v) for v in obj)
return obj
def safe_save(obj: Any, f: str | os.PathLike | BinaryIO, **kwargs) -> None:
"""Save a checkpoint after sanitizing known types for ``weights_only=True`` compatibility."""
torch.save(_sanitize_for_save(obj), f, **kwargs)
def safe_load(f: str | os.PathLike | BinaryIO | bytes, **kwargs) -> Any:
"""Load a checkpoint securely using weights_only=True by default."""
kwargs.setdefault("weights_only", True)
if isinstance(f, (bytes, bytearray)):
f = BytesIO(f)
return torch.load(f, **kwargs)
# Add safe globals for serialization
torch.serialization.add_safe_globals([slice])
+1 -3
View File
@@ -48,9 +48,7 @@ def _get_speech_dataset(dataset_name: str, num_samples: int):
# Use streaming can reduce the downloading time for large datasets
dataset = load_dataset(
**SUPPORTED_SPEECH_DATASET_CONFIG[dataset_name]["config"],
trust_remote_code=True,
streaming=True,
**SUPPORTED_SPEECH_DATASET_CONFIG[dataset_name]["config"], streaming=True
)
else:
raise NotImplementedError(
@@ -22,6 +22,7 @@ from transformers import AutoModelForCausalLM
import modelopt.torch.quantization as mtq
from modelopt.torch.export import export_hf_vllm_fq_checkpoint
from modelopt.torch.quantization.model_quant import fold_weight
from modelopt.torch.utils import safe_load
@pytest.mark.parametrize("quant_cfg", [mtq.FP8_DEFAULT_CFG])
@@ -99,8 +100,7 @@ def test_hf_vllm_export(tmp_path, quant_cfg):
)
# Verify quantizer state dict: same keys, weight quantizer amaxes cleared, input amaxes kept
# weights_only=False required: modelopt_state contains Python objects (dicts, strings, etc.)
quantizer_state_dict = torch.load(modelopt_state_file)["modelopt_state_weights"]
quantizer_state_dict = safe_load(modelopt_state_file)["modelopt_state_weights"]
assert len(quantizer_state_dict) > 0, (
f"modelopt_state_weights should not be empty in {modelopt_state_file}"
)
+1 -3
View File
@@ -163,9 +163,7 @@ def test_gptq_e2e_flow(quant_cfg):
model = AutoModelForCausalLM.from_pretrained(
"TinyLlama/TinyLlama-1.1B-Chat-v1.0", device_map="auto"
)
tokenizer = AutoTokenizer.from_pretrained(
"TinyLlama/TinyLlama-1.1B-Chat-v1.0", trust_remote_code=True
)
tokenizer = AutoTokenizer.from_pretrained("TinyLlama/TinyLlama-1.1B-Chat-v1.0")
# can't set attribute 'pad_token' for "<unk>"
# We skip this step for Nemo models
@@ -29,6 +29,7 @@ from modelopt.torch.quantization.algorithms import (
estimate_quant_compression,
)
from modelopt.torch.quantization.config import _base_disable_all, _default_disabled_quantizer_cfg
from modelopt.torch.utils import safe_load
from modelopt.torch.utils.distributed import DistributedProcessGroup
@@ -421,7 +422,7 @@ def test_auto_quantize_checkpoint_resume(method, tmp_path, capsys):
)
# Verify method is correctly persisted in checkpoint and state dicts
saved = torch.load(checkpoint_path, weights_only=False)
saved = safe_load(checkpoint_path)
assert saved["method"] == method
assert state_dict_1["method"] == method
assert state_dict_2["method"] == method
@@ -0,0 +1,72 @@
# SPDX-FileCopyrightText: Copyright (c) 2023-2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Tests for Modelopt's serialization utilities."""
from io import BytesIO
import torch
from modelopt.torch.opt.config import ModeloptBaseConfig
from modelopt.torch.utils import safe_load
class MockConfig(ModeloptBaseConfig):
"""A mock configuration class for testing serialization."""
name: str = "mock"
def test_safe_load_with_modelopt_config():
"""Verify that safe_load can handle ModeloptBaseConfig subclasses with weights_only=True."""
config = MockConfig(name="test_serialization")
state = {"config": config}
buffer = BytesIO()
torch.save(state, buffer)
data = buffer.getvalue()
# safe_load defaults to weights_only=True
loaded_state = safe_load(data)
assert isinstance(loaded_state["config"], MockConfig)
assert loaded_state["config"].name == "test_serialization"
def test_safe_load_basic_types():
"""Verify that safe_load can handle basic types (standard torch.load functionality)."""
state = {"t": torch.ones(2), "v": [1, 2, 3], "d": {"a": 1}}
buffer = BytesIO()
torch.save(state, buffer)
data = buffer.getvalue()
loaded_state = safe_load(data)
assert torch.allclose(loaded_state["t"], torch.ones(2))
assert loaded_state["v"] == [1, 2, 3]
assert loaded_state["d"]["a"] == 1
def test_safe_load_with_path(tmp_path):
"""Verify that safe_load can handle file paths."""
state = {"data": 42}
file_path = tmp_path / "test.pt"
torch.save(state, file_path)
loaded_state = safe_load(file_path)
assert loaded_state["data"] == 42