已合并
fix(benchmark): restrict remote model code execution #46981
fix(benchmark): restrict remote model code execution #46981
已合并
liuyutong创建于 16 天前
共 10 个文件变更+35-24
@@ -34,6 +34,10 @@ pip install -r ../utils/requirements.txt
34python ../utils/download_hf.py --model baichuan-inc/Baichuan2-7B-Chat --save_path ./Baichuan2-7B-Chat34python ../utils/download_hf.py --model baichuan-inc/Baichuan2-7B-Chat --save_path ./Baichuan2-7B-Chat
35```35```
36 36 
37+> [!WARNING]
38+> Baichuan2 依赖模型仓中的自定义 Python 代码。脚本默认不执行远程代码;仅在确认模型来源可信后使用 `--trust-remote-code`。
39+> 使用远程模型仓时必须同时通过 `--revision` 指定完整的 40 位 commit hash,避免已审核代码被后续更新替换。优先使用本地已校验的模型目录。
40+ 
37## 2. 数据获取41## 2. 数据获取
38 42 
39训练数据集来自 LlamaFactory 仓库示例数据:43训练数据集来自 LlamaFactory 仓库示例数据:
@@ -20,4 +20,5 @@ python train_baichuan2_7B.py \
20 --use_lora \20 --use_lora \
21 --use_bf16 \21 --use_bf16 \
22 --gradient_checkpointing \22 --gradient_checkpointing \
23- > logs/train_baichuan.log 2>&123+ "$@" \
24+ > logs/train_baichuan.log 2>&1
@@ -46,7 +46,8 @@ class Baichuan2Trainer:
46 46 
47 self.tokenizer = AutoTokenizer.from_pretrained(47 self.tokenizer = AutoTokenizer.from_pretrained(
48 self.args.model_path,48 self.args.model_path,
49- trust_remote_code=True,49+ revision=self.args.revision,
50+ trust_remote_code=self.args.trust_remote_code,
50 padding_side="right",51 padding_side="right",
51 model_max_length=self.args.max_length,52 model_max_length=self.args.max_length,
52 )53 )
@@ -56,7 +57,8 @@ class Baichuan2Trainer:
56 57 
57 self.model = AutoModelForCausalLM.from_pretrained(58 self.model = AutoModelForCausalLM.from_pretrained(
58 self.args.model_path,59 self.args.model_path,
59- trust_remote_code=True,60+ revision=self.args.revision,
61+ trust_remote_code=self.args.trust_remote_code,
60 torch_dtype=torch.bfloat16 if self.args.use_bf16 else torch.float16,62 torch_dtype=torch.bfloat16 if self.args.use_bf16 else torch.float16,
61 use_cache=not self.args.gradient_checkpointing,63 use_cache=not self.args.gradient_checkpointing,
62 )64 )
@@ -238,6 +240,10 @@ def build_argparser():
238 parser.add_argument("--mfusion", action="store_true",240 parser.add_argument("--mfusion", action="store_true",
239 help="Enable MFusion for graph fusion optimization")241 help="Enable MFusion for graph fusion optimization")
240 parser.add_argument("--model_path", type=str, required=True)242 parser.add_argument("--model_path", type=str, required=True)
243+ parser.add_argument("--trust-remote-code", action="store_true",
244+ help="Allow execution of custom code from the model repository")
245+ parser.add_argument("--revision", type=str, default=None,
246+ help="Full commit hash for a remote model repository")
241 parser.add_argument("--data_path", type=str, required=True)247 parser.add_argument("--data_path", type=str, required=True)
242 parser.add_argument("--output_dir", type=str, default="./baichuan2-finetuned")248 parser.add_argument("--output_dir", type=str, default="./baichuan2-finetuned")
243 parser.add_argument("--num_epochs", type=int, default=3)249 parser.add_argument("--num_epochs", type=int, default=3)
@@ -278,7 +284,17 @@ def build_argparser():
278 284 
279 285 
280def main():286def main():
281- args = build_argparser().parse_args()287+ parser = build_argparser()
288+ args = parser.parse_args()
289+ if args.trust_remote_code and not os.path.isdir(args.model_path):
290+ revision = args.revision or ""
291+ is_commit_hash = len(revision) == 40 and all(
292+ char in "0123456789abcdefABCDEF" for char in revision
293+ )
294+ if not is_commit_hash:
295+ parser.error(
296+ "--revision must be a full commit hash when enabling remote code"
297+ )
282 detect_device_type()298 detect_device_type()
283 os.environ['TORCHINDUCTOR_NPU_BACKEND']=args.npu_backend299 os.environ['TORCHINDUCTOR_NPU_BACKEND']=args.npu_backend
284 if args.npu_backend == "akg":300 if args.npu_backend == "akg":
@@ -305,4 +321,4 @@ def main():
305 print(f"op_compile_time:{op_compile_time * 1e3} ms", )321 print(f"op_compile_time:{op_compile_time * 1e3} ms", )
306 322 
307if __name__ == "__main__":323if __name__ == "__main__":
308- main()324+ main()
@@ -23,6 +23,9 @@
23 23 
24可在链接页面中 `Files and versions` 一栏直接下载。24可在链接页面中 `Files and versions` 一栏直接下载。
25 25 
26+> [!NOTE]
27+> 本脚本使用 Transformers 内置的 GLM4 实现,不执行模型仓自定义 Python 代码。请使用与 `transformers==4.57.1` 兼容的原生 GLM4 模型。
28+ 
26### 环境提示29### 环境提示
27 30 
28本项目依赖已整理到 `requirements.txt`,可直接安装:31本项目依赖已整理到 `requirements.txt`,可直接安装:
@@ -72,7 +72,6 @@ class GLM4Trainer:
72 self.args.model_path,72 self.args.model_path,
73 quantization_config=bnb_config if self.args.use_4bit else None,73 quantization_config=bnb_config if self.args.use_4bit else None,
74 torch_dtype=torch.bfloat16 if self.args.use_bf16 else torch.float16,74 torch_dtype=torch.bfloat16 if self.args.use_bf16 else torch.float16,
75- trust_remote_code=True,
76 )75 )
77 76 
78 logger.info(f"Moving model to {self.args.device_type}...")77 logger.info(f"Moving model to {self.args.device_type}...")
@@ -80,7 +79,6 @@ class GLM4Trainer:
80 79 
81 self.tokenizer = AutoTokenizer.from_pretrained(80 self.tokenizer = AutoTokenizer.from_pretrained(
82 self.args.model_path,81 self.args.model_path,
83- trust_remote_code=True
84 )82 )
85 83 
86 if self.tokenizer.pad_token is None:84 if self.tokenizer.pad_token is None:
@@ -45,7 +45,6 @@ class GPT_OSS_20BTrainer:
45 self.model = AutoModelForCausalLM.from_pretrained(45 self.model = AutoModelForCausalLM.from_pretrained(
46 self.args.model_path,46 self.args.model_path,
47 torch_dtype=torch.bfloat16 if self.args.use_bf16 else torch.float16,47 torch_dtype=torch.bfloat16 if self.args.use_bf16 else torch.float16,
48- trust_remote_code=True
49 )48 )
50 49 
51 logger.info(f"Moving model to {self.args.device_type}...")50 logger.info(f"Moving model to {self.args.device_type}...")
@@ -53,7 +52,6 @@ class GPT_OSS_20BTrainer:
53 52 
54 self.tokenizer = AutoTokenizer.from_pretrained(53 self.tokenizer = AutoTokenizer.from_pretrained(
55 self.args.model_path,54 self.args.model_path,
56- trust_remote_code=True
57 )55 )
58 56 
59 if self.tokenizer.pad_token is None:57 if self.tokenizer.pad_token is None:
@@ -390,4 +388,4 @@ def main():
390 print("="*50)388 print("="*50)
391 389 
392if __name__ == "__main__":390if __name__ == "__main__":
393- main()391+ main()
@@ -56,14 +56,12 @@ class LLama3Trainer:
56 self.args.model_path,56 self.args.model_path,
57 quantization_config=bnb_config if self.args.use_4bit else None,57 quantization_config=bnb_config if self.args.use_4bit else None,
58 torch_dtype=torch.bfloat16 if self.args.use_bf16 else torch.float32,58 torch_dtype=torch.bfloat16 if self.args.use_bf16 else torch.float32,
59- trust_remote_code=True
60 )59 )
61 60 
62 logger.info(f"Moving model to {self.args.device_type}...")61 logger.info(f"Moving model to {self.args.device_type}...")
63 self.model = self.model.to(self.args.device_type)62 self.model = self.model.to(self.args.device_type)
64 self.tokenizer = AutoTokenizer.from_pretrained(63 self.tokenizer = AutoTokenizer.from_pretrained(
65 self.args.model_path,64 self.args.model_path,
66- trust_remote_code=True
67 )65 )
68 66 
69 if self.tokenizer.pad_token is None:67 if self.tokenizer.pad_token is None:
@@ -58,7 +58,7 @@ def setup_environment():
58def parse_args():58def parse_args():
59 parser = argparse.ArgumentParser(description="Mamba Codestral LoRA finetune")59 parser = argparse.ArgumentParser(description="Mamba Codestral LoRA finetune")
60 60 
61- parser.add_argument("--model_path", type=str, default="/home/zhangyican/workspace/q4_data/Mamba-Codestral-7B-v0.1", help="Base model path")61+ parser.add_argument("--model_path", type=str, required=True, help="Base model path")
62 parser.add_argument("--data_file", type=str, default="./c4_demo.jsonl", help="Training data file path")62 parser.add_argument("--data_file", type=str, default="./c4_demo.jsonl", help="Training data file path")
63 parser.add_argument("--output_dir", type=str, default="./mamba_codestral_lora_no_quant", help="Output directory")63 parser.add_argument("--output_dir", type=str, default="./mamba_codestral_lora_no_quant", help="Output directory")
64 parser.add_argument("--epochs", type=int, default=3, help="Number of training epochs")64 parser.add_argument("--epochs", type=int, default=3, help="Number of training epochs")
@@ -103,7 +103,6 @@ def main():
103 print(f"Loading tokenizer from: {args.model_path}")103 print(f"Loading tokenizer from: {args.model_path}")
104 tokenizer = AutoTokenizer.from_pretrained(104 tokenizer = AutoTokenizer.from_pretrained(
105 args.model_path,105 args.model_path,
106- trust_remote_code=True,
107 use_fast=True106 use_fast=True
108 )107 )
109 if tokenizer.pad_token is None:108 if tokenizer.pad_token is None:
@@ -113,7 +112,6 @@ def main():
113 base_model = AutoModelForCausalLM.from_pretrained(112 base_model = AutoModelForCausalLM.from_pretrained(
114 args.model_path,113 args.model_path,
115 torch_dtype=torch.bfloat16, # Use bfloat16 precision114 torch_dtype=torch.bfloat16, # Use bfloat16 precision
116- trust_remote_code=True,
117 use_cache=False, # Disable cache to save memory115 use_cache=False, # Disable cache to save memory
118 )116 )
119 117 
@@ -221,4 +219,4 @@ def main():
221 print("Training complete!")219 print("Training complete!")
222 220 
223if __name__ == "__main__":221if __name__ == "__main__":
224- main()222+ main()
@@ -69,8 +69,7 @@ def patch_remove_ops_from_generate_list(op_names=None):
69 69 
70def parse_args():70def parse_args():
71 p = argparse.ArgumentParser(description="Qwen2-VL 2B Single GPU Training")71 p = argparse.ArgumentParser(description="Qwen2-VL 2B Single GPU Training")
72- p.add_argument("--model_name_or_path", type=str,72+ p.add_argument("--model_name_or_path", type=str, required=True)
73- default="/data/zyc/Qwen2-VL-2B-Instruct")
74 p.add_argument("--use_lora", type=str2bool, default=False)73 p.add_argument("--use_lora", type=str2bool, default=False)
75 p.add_argument("--lora_r", type=int, default=64)74 p.add_argument("--lora_r", type=int, default=64)
76 p.add_argument("--lora_alpha", type=int, default=16)75 p.add_argument("--lora_alpha", type=int, default=16)
@@ -247,7 +246,6 @@ def main():
247 args.model_name_or_path,246 args.model_name_or_path,
248 min_pixels=args.min_pixels,247 min_pixels=args.min_pixels,
249 max_pixels=args.max_pixels,248 max_pixels=args.max_pixels,
250- trust_remote_code=True,
251 )249 )
252 print(f"Processor loading time: {time.time() - t0:.2f}s")250 print(f"Processor loading time: {time.time() - t0:.2f}s")
253 print(f" min_pixels = {args.min_pixels}, max_pixels = {args.max_pixels}")251 print(f" min_pixels = {args.min_pixels}, max_pixels = {args.max_pixels}")
@@ -260,7 +258,6 @@ def main():
260 model = Qwen2VLForConditionalGeneration.from_pretrained(258 model = Qwen2VLForConditionalGeneration.from_pretrained(
261 args.model_name_or_path,259 args.model_name_or_path,
262 torch_dtype=torch.bfloat16,260 torch_dtype=torch.bfloat16,
263- trust_remote_code=True,
264 )261 )
265 print(f"Model loading time: {time.time() - t0:.2f}s")262 print(f"Model loading time: {time.time() - t0:.2f}s")
266 263 
@@ -368,4 +365,4 @@ def main():
368 processor.save_pretrained(final_output_dir)365 processor.save_pretrained(final_output_dir)
369 366 
370if __name__ == "__main__":367if __name__ == "__main__":
371- main()368+ main()
@@ -57,7 +57,6 @@ class Qwen3Trainer:
57 self.args.model_path,57 self.args.model_path,
58 quantization_config=bnb_config if self.args.use_4bit else None,58 quantization_config=bnb_config if self.args.use_4bit else None,
59 torch_dtype=torch.bfloat16 if self.args.use_bf16 else torch.float32,59 torch_dtype=torch.bfloat16 if self.args.use_bf16 else torch.float32,
60- trust_remote_code=True
61 )60 )
62 61 
63 logger.info(f"Moving model to {self.args.device_type}...")62 logger.info(f"Moving model to {self.args.device_type}...")
@@ -65,7 +64,6 @@ class Qwen3Trainer:
65 64 
66 self.tokenizer = AutoTokenizer.from_pretrained(65 self.tokenizer = AutoTokenizer.from_pretrained(
67 self.args.model_path,66 self.args.model_path,
68- trust_remote_code=True
69 )67 )
70 68 
71 if self.tokenizer.pad_token is None:69 if self.tokenizer.pad_token is None:
@@ -380,4 +378,4 @@ def main():
380 print("="*50)378 print("="*50)
381 379 
382if __name__ == "__main__":380if __name__ == "__main__":
383- main()381+ main()