From 179dcc2a2b67f0fa69273be777dbb02d9845bbc7 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E8=B0=A2=E7=BF=8A=E5=87=A1?= Date: Fri, 28 Aug 2026 16:32:19 +0800 Subject: [PATCH] fix: use dtype instead of deprecated torch_dtype for transformers >= 4.56 config.torch_dtype and the torch_dtype keyword argument were deprecated in transformers 4.56 (PR #39782). Pass dtype based on the installed transformers version (packaging.version), falling back to torch_dtype on older versions. --- train.py | 12 ++++++++++-- 1 file changed, 10 insertions(+), 2 deletions(-) diff --git a/train.py b/train.py index 7409ab1..65a4fa3 100644 --- a/train.py +++ b/train.py @@ -29,6 +29,14 @@ from data.dataset import HybridDataset, collate_fn from data.data_utils import AverageMeter, ProgressMeter, Summary, dict_to_cuda from utils.utils import save_args_to_json, create_log_dir +from packaging.version import Version + +def _dtype_kwargs(dtype): + """`dtype` keyword of `from_pretrained` exists since transformers 4.56 (PR #39782); + older versions use `torch_dtype`.""" + if Version(transformers.__version__) >= Version("4.56"): + return {"dtype": dtype} + return {"torch_dtype": dtype} def env_init(distributed=True): print("Init Env for Distributed Training") @@ -330,7 +338,7 @@ def main(args): model = ShowUIForConditionalGeneration.from_pretrained( model_url, - torch_dtype=torch_dtype, + **_dtype_kwargs(torch_dtype), low_cpu_mem_usage=True, _attn_implementation=args.attn_imple, quantization_config=bnb_config, @@ -343,7 +351,7 @@ def main(args): model = Qwen2VLForConditionalGeneration.from_pretrained( model_url, - torch_dtype=torch_dtype, + **_dtype_kwargs(torch_dtype), low_cpu_mem_usage=True, _attn_implementation=args.attn_imple, quantization_config=bnb_config,