Compare commits

...
1 Commits
Author SHA1 Message Date
SolitaryThinker b30d98ca14 fp8 2025-12-31 00:19:52 +00:00
6 changed files with 157 additions and 7 deletions
+1
View File
@@ -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.)
+1 -1
View File
@@ -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
+17 -1
View File
@@ -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)})"
+38
View File
@@ -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
+73 -3
View File
@@ -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):
+27 -2
View File
@@ -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)