Compare commits
1
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
b30d98ca14 |
@@ -75,3 +75,4 @@ if __name__ == '__main__':
|
||||
- [Configuration](../inference/configuration.md) - Learn about configuration options
|
||||
- [Examples](../inference/examples/) - Explore more examples
|
||||
- [Optimizations](../inference/optimizations.md) - Performance optimization tips
|
||||
- [Low VRAM Inference](../inference/low_vram_inference.md) - Memory-saving settings (CPU offload, sharded loading, etc.)
|
||||
|
||||
@@ -98,7 +98,7 @@ Common issues and their solutions:
|
||||
### Out of Memory Errors
|
||||
If you encounter CUDA out of memory errors:
|
||||
- Reduce `num_frames` or video resolution
|
||||
- Enable memory optimization with `enable_model_cpu_offload`
|
||||
- Enable memory optimization with CPU-offload and sharded loading flags (see [Low VRAM Inference](low_vram_inference.md))
|
||||
- Try a smaller model or use distilled versions
|
||||
- Use `num_gpus` > 1 if multiple GPUs are available
|
||||
|
||||
|
||||
@@ -12,7 +12,8 @@ from fastvideo.configs.models import (DiTConfig, EncoderConfig, ModelConfig,
|
||||
from fastvideo.configs.models.encoders import BaseEncoderOutput
|
||||
from fastvideo.configs.utils import update_config_from_args
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.utils import FlexibleArgumentParser, StoreBoolean, shallow_asdict
|
||||
from fastvideo.utils import (FlexibleArgumentParser, StoreBoolean,
|
||||
PRECISION_TO_TYPE, shallow_asdict)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -311,6 +312,21 @@ class PipelineConfig:
|
||||
f"Length of text encoder configs ({len(self.text_encoder_configs)}) must be equal to length of text encoder precisions ({len(self.text_encoder_precisions)})"
|
||||
)
|
||||
|
||||
unsupported_precisions = [
|
||||
precision for precision in self.text_encoder_precisions
|
||||
if precision not in PRECISION_TO_TYPE
|
||||
and not precision.startswith("fp8")
|
||||
]
|
||||
if unsupported_precisions:
|
||||
supported = ", ".join(PRECISION_TO_TYPE.keys())
|
||||
logger.warning(
|
||||
"Unsupported text encoder precision(s) detected in config: %s. "
|
||||
"FastVideo will attempt to load them with transformers AutoModel when possible. "
|
||||
"Supported fast paths: %s.",
|
||||
unsupported_precisions,
|
||||
supported,
|
||||
)
|
||||
|
||||
if len(self.text_encoder_configs) != len(self.preprocess_text_funcs):
|
||||
raise ValueError(
|
||||
f"Length of text encoder configs ({len(self.text_encoder_configs)}) must be equal to length of text preprocessing functions ({len(self.preprocess_text_funcs)})"
|
||||
|
||||
@@ -170,6 +170,9 @@ class FastVideoArgs:
|
||||
"transformer": True,
|
||||
"vae": True,
|
||||
})
|
||||
text_encoder_override: str | None = None
|
||||
text_encoder_override_path: str | None = None
|
||||
text_encoder_dtype: str | None = None
|
||||
override_transformer_cls_name: str | None = None
|
||||
init_weights_from_safetensors: str = "" # path to safetensors file for initial weight loading
|
||||
init_weights_from_safetensors_2: str = "" # path to safetensors file for initial weight loading for transformer_2
|
||||
@@ -428,6 +431,31 @@ class FastVideoArgs:
|
||||
help=
|
||||
"Use CPU offload for text encoder. Enable if run out of memory.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--text-encoder-override",
|
||||
type=str,
|
||||
default=FastVideoArgs.text_encoder_override,
|
||||
help=
|
||||
("Load text encoder weights from a different local path or HF repo "
|
||||
"instead of the main diffusers snapshot."),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--text-encoder-override-path",
|
||||
type=str,
|
||||
default=FastVideoArgs.text_encoder_override_path,
|
||||
help=
|
||||
("Optional relative path to the text encoder inside the override repository "
|
||||
"(defaults to the module name)."),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--text-encoder-dtype",
|
||||
type=str,
|
||||
default=FastVideoArgs.text_encoder_dtype,
|
||||
help=
|
||||
("Torch dtype string for loading text encoders with transformers AutoModel "
|
||||
"(e.g., fp16, bf16, fp32, fp8). If set, overrides pipeline-config precisions."
|
||||
),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--image-encoder-cpu-offload",
|
||||
action=StoreBoolean,
|
||||
@@ -579,6 +607,16 @@ class FastVideoArgs:
|
||||
kwargs['preprocess_config'] = PreprocessConfig.from_kwargs(kwargs)
|
||||
return cls(**kwargs)
|
||||
|
||||
def get_component_override(
|
||||
self, module_name: str) -> tuple[str | None, str | None]:
|
||||
"""Return override repo/path for a given module if configured."""
|
||||
|
||||
if module_name.startswith(
|
||||
"text_encoder") and self.text_encoder_override:
|
||||
return self.text_encoder_override, self.text_encoder_override_path
|
||||
|
||||
return None, None
|
||||
|
||||
def check_fastvideo_args(self) -> None:
|
||||
"""Validate inference arguments for consistency"""
|
||||
from fastvideo.platforms import current_platform
|
||||
|
||||
@@ -15,7 +15,7 @@ import torch.distributed as dist
|
||||
import torch.nn as nn
|
||||
from safetensors.torch import load_file as safetensors_load_file
|
||||
from torch.distributed import init_device_mesh
|
||||
from transformers import AutoImageProcessor, AutoTokenizer
|
||||
from transformers import AutoImageProcessor, AutoModel, AutoTokenizer
|
||||
from transformers import UMT5EncoderModel
|
||||
from transformers.utils import SAFE_WEIGHTS_INDEX_NAME
|
||||
|
||||
@@ -237,10 +237,23 @@ class TextEncoderLoader(ComponentLoader):
|
||||
encoder_precision = fastvideo_args.pipeline_config.text_encoder_precisions[
|
||||
1]
|
||||
|
||||
requested_dtype = fastvideo_args.text_encoder_dtype or encoder_precision
|
||||
|
||||
target_device = get_local_torch_device()
|
||||
# TODO(will): add support for other dtypes
|
||||
if requested_dtype not in PRECISION_TO_TYPE:
|
||||
logger.info(
|
||||
"Loading text encoder via transformers AutoModel with dtype=%s",
|
||||
requested_dtype,
|
||||
)
|
||||
return self._load_with_transformers(
|
||||
model_path,
|
||||
requested_dtype,
|
||||
fastvideo_args,
|
||||
target_device,
|
||||
)
|
||||
|
||||
return self.load_model(model_path, encoder_config, target_device,
|
||||
fastvideo_args, encoder_precision)
|
||||
fastvideo_args, requested_dtype)
|
||||
|
||||
def load_model(self,
|
||||
model_path: str,
|
||||
@@ -257,6 +270,18 @@ class TextEncoderLoader(ComponentLoader):
|
||||
target_device = torch.device(
|
||||
"mps") if current_platform.is_mps() else torch.device("cpu")
|
||||
|
||||
if dtype not in PRECISION_TO_TYPE:
|
||||
supported = ", ".join(PRECISION_TO_TYPE.keys())
|
||||
raise ValueError(
|
||||
f"Unsupported text encoder precision '{dtype}'. "
|
||||
f"Supported precisions: {supported}. "
|
||||
"FP8 checkpoints will currently be materialized in a supported dtype, "
|
||||
"so they will not reduce VRAM usage."
|
||||
)
|
||||
|
||||
logger.info("Loading text encoder with precision=%s (%s)",
|
||||
dtype, PRECISION_TO_TYPE[dtype])
|
||||
|
||||
with set_default_torch_dtype(PRECISION_TO_TYPE[dtype]):
|
||||
with target_device:
|
||||
architectures = getattr(model_config, "architectures", [])
|
||||
@@ -320,6 +345,51 @@ class TextEncoderLoader(ComponentLoader):
|
||||
|
||||
return model.eval()
|
||||
|
||||
def _resolve_torch_dtype(self, dtype: str) -> torch.dtype:
|
||||
if dtype in PRECISION_TO_TYPE:
|
||||
return PRECISION_TO_TYPE[dtype]
|
||||
|
||||
if dtype.startswith("fp8"):
|
||||
torch_dtype = getattr(torch, "float8_e4m3fn", None)
|
||||
if torch_dtype is None:
|
||||
torch_dtype = getattr(torch, "float8_e4m3fnuz", None)
|
||||
|
||||
if torch_dtype is None:
|
||||
raise ValueError(
|
||||
"FP8 requested for text encoder loading, but the current "
|
||||
"PyTorch build does not expose float8 dtypes. Upgrade PyTorch "
|
||||
"or choose a supported dtype (fp16/bf16/fp32)."
|
||||
)
|
||||
return torch_dtype
|
||||
|
||||
raise ValueError(
|
||||
f"Unsupported text encoder dtype '{dtype}'. "
|
||||
"Pass a torch dtype string such as fp16, bf16, fp32, or fp8."
|
||||
)
|
||||
|
||||
def _load_with_transformers(
|
||||
self,
|
||||
model_path: str,
|
||||
dtype: str,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
target_device: torch.device,
|
||||
) -> nn.Module:
|
||||
torch_dtype = self._resolve_torch_dtype(dtype)
|
||||
device_map = "auto" if fastvideo_args.text_encoder_cpu_offload else None
|
||||
|
||||
model = AutoModel.from_pretrained(
|
||||
model_path,
|
||||
trust_remote_code=fastvideo_args.trust_remote_code,
|
||||
revision=fastvideo_args.revision,
|
||||
torch_dtype=torch_dtype,
|
||||
device_map=device_map,
|
||||
)
|
||||
|
||||
if device_map is None:
|
||||
model = model.to(target_device)
|
||||
|
||||
return model.eval()
|
||||
|
||||
|
||||
class ImageEncoderLoader(TextEncoderLoader):
|
||||
|
||||
|
||||
@@ -352,8 +352,8 @@ class ComposedPipelineBase(ABC):
|
||||
else:
|
||||
load_module_name = module_name
|
||||
|
||||
component_model_path = os.path.join(self.model_path,
|
||||
load_module_name)
|
||||
component_model_path = self._resolve_component_model_path(
|
||||
fastvideo_args, module_name, load_module_name)
|
||||
module = PipelineComponentLoader.load_module(
|
||||
module_name=load_module_name,
|
||||
component_model_path=component_model_path,
|
||||
@@ -376,6 +376,31 @@ class ComposedPipelineBase(ABC):
|
||||
|
||||
return modules
|
||||
|
||||
def _resolve_component_model_path(
|
||||
self,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
module_name: str,
|
||||
load_module_name: str,
|
||||
) -> str:
|
||||
"""Resolve the on-disk path for a given module, respecting overrides."""
|
||||
|
||||
override_repo, override_subpath = fastvideo_args.get_component_override(
|
||||
module_name)
|
||||
if override_repo is not None:
|
||||
override_root = maybe_download_model(override_repo)
|
||||
module_subpath = override_subpath or load_module_name
|
||||
override_path = os.path.join(override_root, module_subpath)
|
||||
if not os.path.exists(override_path):
|
||||
raise FileNotFoundError(
|
||||
f"Override path for {module_name} not found: {override_path}"
|
||||
)
|
||||
|
||||
logger.info("Using override for %s from %s", module_name,
|
||||
override_path)
|
||||
return override_path
|
||||
|
||||
return os.path.join(self.model_path, load_module_name)
|
||||
|
||||
def add_stage(self, stage_name: str, stage: PipelineStage):
|
||||
assert self.modules is not None, "No modules are registered"
|
||||
self._stages.append(stage)
|
||||
|
||||
Reference in New Issue
Block a user