From 43739895a18ae11fe66f45690d603aeb57ba360f Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E8=B0=A2=E7=BF=8A=E5=87=A1?= <165987053+xyf5432@users.noreply.github.com> Date: Wed, 2 Sep 2026 16:47:39 +0800 Subject: [PATCH] fix: use dtype instead of deprecated torch_dtype (#514) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * 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). Use dtype when accepted and fall back to torch_dtype via try/except TypeError for older versions. * 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. --------- Co-authored-by: 谢翊凡 --- .../z_image_fun/predict_t2i_control_2.1_lite.py | 13 +++++++++++-- .../z_image_fun/predict_turbo_t2i_control_2.1.py | 13 +++++++++++-- 2 files changed, 22 insertions(+), 4 deletions(-) diff --git a/examples/z_image_fun/predict_t2i_control_2.1_lite.py b/examples/z_image_fun/predict_t2i_control_2.1_lite.py index 6defa3c..47f07ac 100644 --- a/examples/z_image_fun/predict_t2i_control_2.1_lite.py +++ b/examples/z_image_fun/predict_t2i_control_2.1_lite.py @@ -28,6 +28,15 @@ from videox_fun.utils.utils import (filter_kwargs, get_image, get_image_latent, get_image_to_video_latent, get_video_to_video_latent, save_videos_grid) +import transformers +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} # GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, sequential_cpu_offload]. # model_full_load means that the entire model will be moved to the GPU. @@ -138,7 +147,7 @@ tokenizer = AutoTokenizer.from_pretrained( model_name, subfolder="tokenizer" ) text_encoder = Qwen3ForCausalLM.from_pretrained( - model_name, subfolder="text_encoder", torch_dtype=weight_dtype, + model_name, subfolder="text_encoder", **_dtype_kwargs(weight_dtype), low_cpu_mem_usage=True, ) @@ -247,4 +256,4 @@ if ulysses_degree * ring_degree > 1: if dist.get_rank() == 0: save_results() else: - save_results() \ No newline at end of file + save_results() diff --git a/examples/z_image_fun/predict_turbo_t2i_control_2.1.py b/examples/z_image_fun/predict_turbo_t2i_control_2.1.py index 7b38bcf..dd833eb 100644 --- a/examples/z_image_fun/predict_turbo_t2i_control_2.1.py +++ b/examples/z_image_fun/predict_turbo_t2i_control_2.1.py @@ -28,6 +28,15 @@ from videox_fun.utils.utils import (filter_kwargs, get_image, get_image_latent, get_image_to_video_latent, get_video_to_video_latent, save_videos_grid) +import transformers +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} # GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, sequential_cpu_offload]. # model_full_load means that the entire model will be moved to the GPU. @@ -138,7 +147,7 @@ tokenizer = AutoTokenizer.from_pretrained( model_name, subfolder="tokenizer" ) text_encoder = Qwen3ForCausalLM.from_pretrained( - model_name, subfolder="text_encoder", torch_dtype=weight_dtype, + model_name, subfolder="text_encoder", **_dtype_kwargs(weight_dtype), low_cpu_mem_usage=True, ) @@ -247,4 +256,4 @@ if ulysses_degree * ring_degree > 1: if dist.get_rank() == 0: save_results() else: - save_results() \ No newline at end of file + save_results()