From 5dc17dfd15d9ce5aa24553c2879ac68a4787f6bb Mon Sep 17 00:00:00 2001 From: Keval Morabia <28916987+kevalmorabia97@users.noreply.github.com> Date: Wed, 8 Apr 2026 00:36:28 +0530 Subject: [PATCH] [Security] Enable torch.load(weights_only=True) for secure checkpoint loading + trust_remote_code fix (#1181) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit ### 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 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 you copied code from any other sources or added a new PIP dependency, did you follow guidance in `CONTRIBUTING.md`: ✅ - Did you write any new necessary tests?: ✅ - Did you update [Changelog](https://github.com/NVIDIA/Model-Optimizer/blob/main/CHANGELOG.rst)?: ✅ ### Additional Information NVBug: 5999336 ## 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. --------- 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> --- CHANGELOG.rst | 4 ++ .../gpt-oss/convert_oai_mxfp4_weight_only.py | 17 +++-- examples/llm_autodeploy/run_auto_quantize.py | 11 ++- examples/llm_eval/lm_eval_hf.py | 3 +- examples/llm_eval/modeling.py | 23 +++--- examples/llm_ptq/example_utils.py | 14 ++-- examples/llm_ptq/hf_ptq.py | 2 + examples/llm_ptq/vlm_utils.py | 22 ++++-- .../llm_qad/data_utils/download_dataset.py | 11 ++- .../compute_hidden_states_hf.py | 9 ++- .../scripts/ar_validate.py | 9 ++- .../scripts/export_hf_checkpoint.py | 10 +-- .../scripts/send_conversation_vllm.py | 9 ++- .../accuracy_benchmark/mmlu_benchmark.py | 6 +- .../windows/accuracy_benchmark/modeling.py | 21 ++++-- examples/windows/onnx_ptq/whisper/README.md | 2 +- .../whisper/whisper_onnx_quantization.py | 2 +- .../whisper/whisper_optimum_ort_inference.py | 4 +- .../sample_example_qad_diffusers.py | 5 +- modelopt/torch/export/distribute.py | 6 +- modelopt/torch/export/model_config.py | 23 ++++++ modelopt/torch/opt/config.py | 9 +++ modelopt/torch/opt/conversion.py | 14 ++-- modelopt/torch/opt/hparam.py | 8 +++ .../opt/plugins/mcore_dist_checkpointing.py | 6 +- modelopt/torch/opt/plugins/megatron.py | 15 ++-- modelopt/torch/opt/plugins/peft.py | 7 +- modelopt/torch/opt/searcher.py | 15 ++-- .../prune/importance_hooks/base_hooks.py | 8 +-- .../compare_module_outputs.py | 17 ++--- .../torch/quantization/calib/calibrator.py | 12 ++++ .../quantization/qtensor/base_qtensor.py | 13 ++++ modelopt/torch/utils/__init__.py | 1 + modelopt/torch/utils/serialization.py | 67 +++++++++++++++++ modelopt/torch/utils/speech_dataset_utils.py | 4 +- .../export/test_vllm_fakequant_hf_export.py | 4 +- tests/gpu/torch/quantization/test_gptq.py | 4 +- .../unit/torch/quantization/test_autoquant.py | 3 +- tests/unit/torch/utils/test_serialization.py | 72 +++++++++++++++++++ 39 files changed, 381 insertions(+), 111 deletions(-) create mode 100644 modelopt/torch/utils/serialization.py create mode 100644 tests/unit/torch/utils/test_serialization.py diff --git a/CHANGELOG.rst b/CHANGELOG.rst index 37e86e8fe..1836d41c3 100755 --- a/CHANGELOG.rst +++ b/CHANGELOG.rst @@ -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 ` and :meth:`safe_load ` API to save and load checkpoints securely. + 0.43 (2026-04-09) ^^^^^^^^^^^^^^^^^ diff --git a/examples/gpt-oss/convert_oai_mxfp4_weight_only.py b/examples/gpt-oss/convert_oai_mxfp4_weight_only.py index bebb91486..8ebcf4779 100644 --- a/examples/gpt-oss/convert_oai_mxfp4_weight_only.py +++ b/examples/gpt-oss/convert_oai_mxfp4_weight_only.py @@ -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) diff --git a/examples/llm_autodeploy/run_auto_quantize.py b/examples/llm_autodeploy/run_auto_quantize.py index 155fa7434..e6ce2429f 100644 --- a/examples/llm_autodeploy/run_auto_quantize.py +++ b/examples/llm_autodeploy/run_auto_quantize.py @@ -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 != "": @@ -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, ) diff --git a/examples/llm_eval/lm_eval_hf.py b/examples/llm_eval/lm_eval_hf.py index 405e8590a..11d736a42 100755 --- a/examples/llm_eval/lm_eval_hf.py +++ b/examples/llm_eval/lm_eval_hf.py @@ -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 diff --git a/examples/llm_eval/modeling.py b/examples/llm_eval/modeling.py index d06d05560..93732e7f6 100644 --- a/examples/llm_eval/modeling.py +++ b/examples/llm_eval/modeling.py @@ -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) diff --git a/examples/llm_ptq/example_utils.py b/examples/llm_ptq/example_utils.py index 8c81607ec..2a851de5c 100755 --- a/examples/llm_ptq/example_utils.py +++ b/examples/llm_ptq/example_utils.py @@ -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 diff --git a/examples/llm_ptq/hf_ptq.py b/examples/llm_ptq/hf_ptq.py index f7a6fa369..bb240ba0c 100755 --- a/examples/llm_ptq/hf_ptq.py +++ b/examples/llm_ptq/hf_ptq.py @@ -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( diff --git a/examples/llm_ptq/vlm_utils.py b/examples/llm_ptq/vlm_utils.py index 9919e405b..abfebbd4f 100644 --- a/examples/llm_ptq/vlm_utils.py +++ b/examples/llm_ptq/vlm_utils.py @@ -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 = [ diff --git a/examples/llm_qad/data_utils/download_dataset.py b/examples/llm_qad/data_utils/download_dataset.py index e3e3d0646..46b23c142 100644 --- a/examples/llm_qad/data_utils/download_dataset.py +++ b/examples/llm_qad/data_utils/download_dataset.py @@ -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" diff --git a/examples/speculative_decoding/collect_hidden_states/compute_hidden_states_hf.py b/examples/speculative_decoding/collect_hidden_states/compute_hidden_states_hf.py index a3d1681c4..b062d833b 100644 --- a/examples/speculative_decoding/collect_hidden_states/compute_hidden_states_hf.py +++ b/examples/speculative_decoding/collect_hidden_states/compute_hidden_states_hf.py @@ -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, "") diff --git a/examples/speculative_decoding/scripts/ar_validate.py b/examples/speculative_decoding/scripts/ar_validate.py index d1bf31a1a..1ad7bec40 100644 --- a/examples/speculative_decoding/scripts/ar_validate.py +++ b/examples/speculative_decoding/scripts/ar_validate.py @@ -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) diff --git a/examples/speculative_decoding/scripts/export_hf_checkpoint.py b/examples/speculative_decoding/scripts/export_hf_checkpoint.py index 925f4b73d..2771ab151 100644 --- a/examples/speculative_decoding/scripts/export_hf_checkpoint.py +++ b/examples/speculative_decoding/scripts/export_hf_checkpoint.py @@ -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}") diff --git a/examples/speculative_decoding/scripts/send_conversation_vllm.py b/examples/speculative_decoding/scripts/send_conversation_vllm.py index 5101b4e6f..470df80be 100644 --- a/examples/speculative_decoding/scripts/send_conversation_vllm.py +++ b/examples/speculative_decoding/scripts/send_conversation_vllm.py @@ -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 diff --git a/examples/windows/accuracy_benchmark/mmlu_benchmark.py b/examples/windows/accuracy_benchmark/mmlu_benchmark.py index 4eb2fd619..54573e642 100644 --- a/examples/windows/accuracy_benchmark/mmlu_benchmark.py +++ b/examples/windows/accuracy_benchmark/mmlu_benchmark.py @@ -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: diff --git a/examples/windows/accuracy_benchmark/modeling.py b/examples/windows/accuracy_benchmark/modeling.py index 273a944c5..f17300be9 100644 --- a/examples/windows/accuracy_benchmark/modeling.py +++ b/examples/windows/accuracy_benchmark/modeling.py @@ -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() diff --git a/examples/windows/onnx_ptq/whisper/README.md b/examples/windows/onnx_ptq/whisper/README.md index 8757aaeb5..82ae78220 100644 --- a/examples/windows/onnx_ptq/whisper/README.md +++ b/examples/windows/onnx_ptq/whisper/README.md @@ -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. diff --git a/examples/windows/onnx_ptq/whisper/whisper_onnx_quantization.py b/examples/windows/onnx_ptq/whisper/whisper_onnx_quantization.py index 7b3e3d319..03d2c4980 100644 --- a/examples/windows/onnx_ptq/whisper/whisper_onnx_quantization.py +++ b/examples/windows/onnx_ptq/whisper/whisper_onnx_quantization.py @@ -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 diff --git a/examples/windows/onnx_ptq/whisper/whisper_optimum_ort_inference.py b/examples/windows/onnx_ptq/whisper/whisper_optimum_ort_inference.py index 52d56fe04..a1f39b8f0 100644 --- a/examples/windows/onnx_ptq/whisper/whisper_optimum_ort_inference.py +++ b/examples/windows/onnx_ptq/whisper/whisper_optimum_ort_inference.py @@ -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 = [] diff --git a/examples/windows/torch_onnx/diffusers/qad_example/sample_example_qad_diffusers.py b/examples/windows/torch_onnx/diffusers/qad_example/sample_example_qad_diffusers.py index 7ed73b72c..162a297ca 100644 --- a/examples/windows/torch_onnx/diffusers/qad_example/sample_example_qad_diffusers.py +++ b/examples/windows/torch_onnx/diffusers/qad_example/sample_example_qad_diffusers.py @@ -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 diff --git a/modelopt/torch/export/distribute.py b/modelopt/torch/export/distribute.py index 4fe7be43e..97fa6e7eb 100644 --- a/modelopt/torch/export/distribute.py +++ b/modelopt/torch/export/distribute.py @@ -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 diff --git a/modelopt/torch/export/model_config.py b/modelopt/torch/export/model_config.py index 4826d0639..dce39767c 100755 --- a/modelopt/torch/export/model_config.py +++ b/modelopt/torch/export/model_config.py @@ -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, + ] +) diff --git a/modelopt/torch/opt/config.py b/modelopt/torch/opt/config.py index 032b9fe6b..62f7b7e16 100644 --- a/modelopt/torch/opt/config.py +++ b/modelopt/torch/opt/config.py @@ -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} diff --git a/modelopt/torch/opt/conversion.py b/modelopt/torch/opt/conversion.py index 6ec7a1729..432bd6abc 100644 --- a/modelopt/torch/opt/conversion.py +++ b/modelopt/torch/opt/conversion.py @@ -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()`. @@ -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"]) diff --git a/modelopt/torch/opt/hparam.py b/modelopt/torch/opt/hparam.py index f4de5a329..27b97732d 100644 --- a/modelopt/torch/opt/hparam.py +++ b/modelopt/torch/opt/hparam.py @@ -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.""" diff --git a/modelopt/torch/opt/plugins/mcore_dist_checkpointing.py b/modelopt/torch/opt/plugins/mcore_dist_checkpointing.py index 16ace511f..aac9dcc4d 100644 --- a/modelopt/torch/opt/plugins/mcore_dist_checkpointing.py +++ b/modelopt/torch/opt/plugins/mcore_dist_checkpointing.py @@ -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"] diff --git a/modelopt/torch/opt/plugins/megatron.py b/modelopt/torch/opt/plugins/megatron.py index 761e8d9a4..b499cb542 100644 --- a/modelopt/torch/opt/plugins/megatron.py +++ b/modelopt/torch/opt/plugins/megatron.py @@ -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.") diff --git a/modelopt/torch/opt/plugins/peft.py b/modelopt/torch/opt/plugins/peft.py index de1218917..256d1702e 100644 --- a/modelopt/torch/opt/plugins/peft.py +++ b/modelopt/torch/opt/plugins/peft.py @@ -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): diff --git a/modelopt/torch/opt/searcher.py b/modelopt/torch/opt/searcher.py index b9e1c1116..386948cb4 100644 --- a/modelopt/torch/opt/searcher.py +++ b/modelopt/torch/opt/searcher.py @@ -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: diff --git a/modelopt/torch/prune/importance_hooks/base_hooks.py b/modelopt/torch/prune/importance_hooks/base_hooks.py index 248e6ec10..22e82c27b 100644 --- a/modelopt/torch/prune/importance_hooks/base_hooks.py +++ b/modelopt/torch/prune/importance_hooks/base_hooks.py @@ -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(): diff --git a/modelopt/torch/prune/importance_hooks/compare_module_outputs.py b/modelopt/torch/prune/importance_hooks/compare_module_outputs.py index e692a518a..0a7ea542b 100644 --- a/modelopt/torch/prune/importance_hooks/compare_module_outputs.py +++ b/modelopt/torch/prune/importance_hooks/compare_module_outputs.py @@ -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) diff --git a/modelopt/torch/quantization/calib/calibrator.py b/modelopt/torch/quantization/calib/calibrator.py index cff7a75e7..633f0070c 100644 --- a/modelopt/torch/quantization/calib/calibrator.py +++ b/modelopt/torch/quantization/calib/calibrator.py @@ -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. diff --git a/modelopt/torch/quantization/qtensor/base_qtensor.py b/modelopt/torch/quantization/qtensor/base_qtensor.py index d5a9a4269..c621e9fab 100644 --- a/modelopt/torch/quantization/qtensor/base_qtensor.py +++ b/modelopt/torch/quantization/qtensor/base_qtensor.py @@ -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]) diff --git a/modelopt/torch/utils/__init__.py b/modelopt/torch/utils/__init__.py index f026e747a..51d02248c 100644 --- a/modelopt/torch/utils/__init__.py +++ b/modelopt/torch/utils/__init__.py @@ -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 * diff --git a/modelopt/torch/utils/serialization.py b/modelopt/torch/utils/serialization.py new file mode 100644 index 000000000..da16f7514 --- /dev/null +++ b/modelopt/torch/utils/serialization.py @@ -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]) diff --git a/modelopt/torch/utils/speech_dataset_utils.py b/modelopt/torch/utils/speech_dataset_utils.py index a71d73773..ef0660175 100644 --- a/modelopt/torch/utils/speech_dataset_utils.py +++ b/modelopt/torch/utils/speech_dataset_utils.py @@ -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( diff --git a/tests/gpu/torch/export/test_vllm_fakequant_hf_export.py b/tests/gpu/torch/export/test_vllm_fakequant_hf_export.py index 8f6071796..8ee71ed45 100644 --- a/tests/gpu/torch/export/test_vllm_fakequant_hf_export.py +++ b/tests/gpu/torch/export/test_vllm_fakequant_hf_export.py @@ -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}" ) diff --git a/tests/gpu/torch/quantization/test_gptq.py b/tests/gpu/torch/quantization/test_gptq.py index 0c60bcd00..d43177cae 100644 --- a/tests/gpu/torch/quantization/test_gptq.py +++ b/tests/gpu/torch/quantization/test_gptq.py @@ -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 "" # We skip this step for Nemo models diff --git a/tests/unit/torch/quantization/test_autoquant.py b/tests/unit/torch/quantization/test_autoquant.py index 5ba9a5110..87ec73291 100644 --- a/tests/unit/torch/quantization/test_autoquant.py +++ b/tests/unit/torch/quantization/test_autoquant.py @@ -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 diff --git a/tests/unit/torch/utils/test_serialization.py b/tests/unit/torch/utils/test_serialization.py new file mode 100644 index 000000000..32851d3a0 --- /dev/null +++ b/tests/unit/torch/utils/test_serialization.py @@ -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