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:
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user