mirror of
https://github.com/NVIDIA/Model-Optimizer.git
synced 2026-10-02 03:14:52 +08:00
[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:
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
]
|
||||
)
|
||||
|
||||
@@ -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}
|
||||
|
||||
@@ -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"])
|
||||
|
||||
@@ -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"]
|
||||
|
||||
|
||||
@@ -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.")
|
||||
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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])
|
||||
|
||||
@@ -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 *
|
||||
|
||||
@@ -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])
|
||||
@@ -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}"
|
||||
)
|
||||
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user