fix: use dtype instead of deprecated torch_dtype (#514)

* 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: 谢翊凡 <xyf5432@users.noreply.github.com>
This commit is contained in:
谢翊凡
2026-09-02 16:47:39 +08:00
committed by GitHub
co-authored by 谢翊凡
parent 6f3fb60dad
commit 43739895a1
2 changed files with 22 additions and 4 deletions
@@ -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()
save_results()
@@ -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()
save_results()