Compare commits
3
Commits
main
...
will/zmq-worker
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
99dff854ad | ||
|
|
4ccf68d71e | ||
|
|
15b798a9d1 |
@@ -1 +1,3 @@
|
||||
from fastvideo.v1.entrypoints.video_generator import VideoGenerator
|
||||
|
||||
__all__ = ["VideoGenerator"]
|
||||
|
||||
@@ -1,8 +1,9 @@
|
||||
from fastvideo.v1.configs.hunyuan import HunyuanConfig, FastHunyuanConfig
|
||||
from fastvideo.v1.configs.wan import WanT2V480PConfig, WanI2V480PConfig
|
||||
from fastvideo.v1.configs.base import BaseConfig, SlidingTileAttnConfig
|
||||
from fastvideo.v1.configs.registry import get_pipeline_config_cls_for_name
|
||||
|
||||
__all__ = [
|
||||
"HunyuanConfig", "FastHunyuanConfig", "BaseConfig", "SlidingTileAttnConfig",
|
||||
"WanT2V480PConfig", "WanI2V480PConfig"
|
||||
"WanT2V480PConfig", "WanI2V480PConfig", "get_pipeline_config_cls_for_name"
|
||||
]
|
||||
|
||||
@@ -48,9 +48,6 @@ class BaseConfig:
|
||||
|
||||
neg_prompt: Optional[str] = None
|
||||
|
||||
# Additional parameters can be added as a dict
|
||||
extra_params: Dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
class SlidingTileAttnConfig(BaseConfig):
|
||||
|
||||
@@ -40,8 +40,8 @@ PIPELINE_FALLBACK_CONFIG: Dict[str, Type[BaseConfig]] = {
|
||||
}
|
||||
|
||||
|
||||
def get_pipeline_config_for_name(
|
||||
pipeline_name_or_path: str) -> Optional[Type[BaseConfig]]:
|
||||
def get_pipeline_config_cls_for_name(
|
||||
pipeline_name_or_path: str) -> Optional[BaseConfig]:
|
||||
"""Get the appropriate config class for specific pretrained weights."""
|
||||
|
||||
if os.path.exists(pipeline_name_or_path):
|
||||
|
||||
@@ -1,5 +1,17 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from fastvideo.v1.distributed.communication_op import *
|
||||
from fastvideo.v1.distributed.parallel_state import *
|
||||
from fastvideo.v1.distributed.parallel_state import (
|
||||
init_distributed_environment,
|
||||
initialize_model_parallel,
|
||||
get_sequence_model_parallel_rank,
|
||||
get_sequence_model_parallel_world_size,
|
||||
)
|
||||
from fastvideo.v1.distributed.utils import *
|
||||
|
||||
__all__ = [
|
||||
"init_distributed_environment",
|
||||
"initialize_model_parallel",
|
||||
"get_sequence_model_parallel_rank",
|
||||
"get_sequence_model_parallel_world_size",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,402 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
VideoGenerator module for FastVideo.
|
||||
|
||||
This module provides a consolidated interface for generating videos using
|
||||
diffusion models.
|
||||
"""
|
||||
|
||||
import os
|
||||
import time
|
||||
from typing import Any, Callable, Dict, List, Optional, Union
|
||||
from dataclasses import asdict
|
||||
|
||||
import imageio
|
||||
import numpy as np
|
||||
import torch
|
||||
import torchvision
|
||||
from einops import rearrange
|
||||
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines import (ForwardBatch)
|
||||
from fastvideo.v1.configs import get_pipeline_config_cls_for_name
|
||||
|
||||
from fastvideo.v1.utils import align_to
|
||||
from fastvideo.v1.worker.executor import Executor
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class VideoGenerator:
|
||||
"""
|
||||
A unified class for generating videos using diffusion models.
|
||||
|
||||
This class provides a simple interface for video generation with rich
|
||||
customization options, similar to popular frameworks like HF Diffusers.
|
||||
"""
|
||||
|
||||
def __init__(self, fastvideo_args: FastVideoArgs,
|
||||
executor_class: type[Executor], log_stats: bool):
|
||||
"""
|
||||
Initialize the video generator.
|
||||
|
||||
Args:
|
||||
pipeline: The pipeline to use for inference
|
||||
fastvideo_args: The inference arguments
|
||||
"""
|
||||
self.fastvideo_args = fastvideo_args
|
||||
self.executor = executor_class(fastvideo_args)
|
||||
|
||||
@classmethod
|
||||
def from_pretrained(cls,
|
||||
model_path: str,
|
||||
device: Optional[str] = None,
|
||||
torch_dtype: Optional[torch.dtype] = None,
|
||||
**kwargs) -> "VideoGenerator":
|
||||
"""
|
||||
Create a video generator from a pretrained model.
|
||||
|
||||
Args:
|
||||
model_path: Path or identifier for the pretrained model
|
||||
device: Device to load the model on (e.g., "cuda", "cuda:0", "cpu")
|
||||
torch_dtype: Data type for model weights (e.g., torch.float16)
|
||||
**kwargs: Additional arguments to customize model loading
|
||||
|
||||
Returns:
|
||||
The created video generator
|
||||
"""
|
||||
|
||||
config_cls = get_pipeline_config_cls_for_name(model_path)
|
||||
config = config_cls()
|
||||
|
||||
if config is None:
|
||||
logger.warning(f"No config found for model {model_path}, using default config")
|
||||
config_args = {}
|
||||
else:
|
||||
config_args = asdict(config)
|
||||
|
||||
# override config_args with kwargs
|
||||
config_args.update(kwargs)
|
||||
|
||||
fastvideo_args = FastVideoArgs(
|
||||
model_path=model_path,
|
||||
device_str=device or "cuda" if torch.cuda.is_available() else "cpu",
|
||||
**config_args)
|
||||
|
||||
|
||||
if torch_dtype is not None:
|
||||
fastvideo_args.dtype = torch_dtype
|
||||
|
||||
return cls.from_fastvideo_args(fastvideo_args)
|
||||
|
||||
@classmethod
|
||||
def from_fastvideo_args(cls,
|
||||
fastvideo_args: FastVideoArgs) -> "VideoGenerator":
|
||||
"""
|
||||
Create a video generator with the specified arguments.
|
||||
|
||||
Args:
|
||||
fastvideo_args: The inference arguments
|
||||
|
||||
Returns:
|
||||
The created video generator
|
||||
"""
|
||||
# Initialize distributed environment if needed
|
||||
# initialize_distributed_and_parallelism(fastvideo_args)
|
||||
|
||||
executor_class = Executor.get_class(fastvideo_args)
|
||||
|
||||
return cls(
|
||||
fastvideo_args=fastvideo_args,
|
||||
executor_class=executor_class,
|
||||
log_stats=False, # TODO: implement
|
||||
)
|
||||
|
||||
def generate_video(
|
||||
self,
|
||||
prompt: str,
|
||||
negative_prompt: Optional[str] = None,
|
||||
output_path: Optional[str] = None,
|
||||
save_video: bool = True,
|
||||
return_frames: bool = False,
|
||||
num_inference_steps: Optional[int] = None,
|
||||
guidance_scale: Optional[float] = None,
|
||||
num_frames: Optional[int] = None,
|
||||
height: Optional[int] = None,
|
||||
width: Optional[int] = None,
|
||||
fps: Optional[int] = None,
|
||||
seed: Optional[int] = None,
|
||||
callback: Optional[Callable[[int, int, torch.Tensor], None]] = None,
|
||||
callback_steps: int = 1,
|
||||
) -> Union[Dict[str, Any], List[np.ndarray]]:
|
||||
"""
|
||||
Generate a video based on the given prompt.
|
||||
|
||||
Args:
|
||||
prompt: The prompt to use for generation
|
||||
negative_prompt: The negative prompt to use (overrides the one in fastvideo_args)
|
||||
output_path: Path to save the video (overrides the one in fastvideo_args)
|
||||
save_video: Whether to save the video to disk
|
||||
return_frames: Whether to return the raw frames
|
||||
num_inference_steps: Number of denoising steps (overrides fastvideo_args)
|
||||
guidance_scale: Classifier-free guidance scale (overrides fastvideo_args)
|
||||
num_frames: Number of frames to generate (overrides fastvideo_args)
|
||||
height: Height of generated video (overrides fastvideo_args)
|
||||
width: Width of generated video (overrides fastvideo_args)
|
||||
fps: Frames per second for saved video (overrides fastvideo_args)
|
||||
seed: Random seed for generation (overrides fastvideo_args)
|
||||
callback: Callback function called after each step
|
||||
callback_steps: Number of steps between each callback
|
||||
|
||||
Returns:
|
||||
Either the output dictionary or the list of frames depending on return_frames
|
||||
"""
|
||||
# Create a copy of inference args to avoid modifying the original
|
||||
fastvideo_args = self.fastvideo_args.copy()
|
||||
|
||||
# Override parameters if provided
|
||||
if negative_prompt is not None:
|
||||
fastvideo_args.neg_prompt = negative_prompt
|
||||
if num_inference_steps is not None:
|
||||
fastvideo_args.num_inference_steps = num_inference_steps
|
||||
if guidance_scale is not None:
|
||||
fastvideo_args.guidance_scale = guidance_scale
|
||||
if num_frames is not None:
|
||||
fastvideo_args.num_frames = num_frames
|
||||
if height is not None:
|
||||
fastvideo_args.height = height
|
||||
if width is not None:
|
||||
fastvideo_args.width = width
|
||||
if fps is not None:
|
||||
fastvideo_args.fps = fps
|
||||
if seed is not None:
|
||||
fastvideo_args.seed = seed
|
||||
|
||||
# Store callback info
|
||||
fastvideo_args.callback = callback
|
||||
fastvideo_args.callback_steps = callback_steps
|
||||
|
||||
# Validate inputs
|
||||
if not isinstance(prompt, str):
|
||||
raise TypeError(
|
||||
f"`prompt` must be a string, but got {type(prompt)}")
|
||||
prompt = prompt.strip()
|
||||
|
||||
# Process negative prompt
|
||||
if fastvideo_args.neg_prompt is not None:
|
||||
fastvideo_args.neg_prompt = fastvideo_args.neg_prompt.strip()
|
||||
|
||||
# Validate dimensions
|
||||
if (fastvideo_args.height <= 0 or fastvideo_args.width <= 0
|
||||
or fastvideo_args.num_frames <= 0):
|
||||
raise ValueError(
|
||||
f"Height, width, and num_frames must be positive integers, got "
|
||||
f"height={fastvideo_args.height}, width={fastvideo_args.width}, "
|
||||
f"num_frames={fastvideo_args.num_frames}")
|
||||
|
||||
if (fastvideo_args.num_frames - 1) % 4 != 0:
|
||||
raise ValueError(
|
||||
f"num_frames-1 must be a multiple of 4, got {fastvideo_args.num_frames}"
|
||||
)
|
||||
|
||||
# Calculate sizes
|
||||
target_height = align_to(fastvideo_args.height, 16)
|
||||
target_width = align_to(fastvideo_args.width, 16)
|
||||
|
||||
# Calculate latent sizes
|
||||
latents_size = [(fastvideo_args.num_frames - 1) // 4 + 1,
|
||||
fastvideo_args.height // 8, fastvideo_args.width // 8]
|
||||
n_tokens = latents_size[0] * latents_size[1] * latents_size[2]
|
||||
|
||||
# Log parameters
|
||||
debug_str = f"""
|
||||
height: {target_height}
|
||||
width: {target_width}
|
||||
video_length: {fastvideo_args.num_frames}
|
||||
prompt: {prompt}
|
||||
neg_prompt: {fastvideo_args.neg_prompt}
|
||||
seed: {fastvideo_args.seed}
|
||||
infer_steps: {fastvideo_args.num_inference_steps}
|
||||
num_videos_per_prompt: {fastvideo_args.num_videos}
|
||||
guidance_scale: {fastvideo_args.guidance_scale}
|
||||
n_tokens: {n_tokens}
|
||||
flow_shift: {fastvideo_args.flow_shift}
|
||||
embedded_guidance_scale: {fastvideo_args.embedded_cfg_scale}"""
|
||||
logger.info(debug_str)
|
||||
|
||||
# Prepare batch
|
||||
device = torch.device(fastvideo_args.device_str)
|
||||
batch = ForwardBatch(
|
||||
prompt=prompt,
|
||||
negative_prompt=fastvideo_args.neg_prompt,
|
||||
num_videos_per_prompt=fastvideo_args.num_videos,
|
||||
height=fastvideo_args.height,
|
||||
width=fastvideo_args.width,
|
||||
num_frames=fastvideo_args.num_frames,
|
||||
num_inference_steps=fastvideo_args.num_inference_steps,
|
||||
guidance_scale=fastvideo_args.guidance_scale,
|
||||
eta=0.0,
|
||||
n_tokens=n_tokens,
|
||||
data_type="video" if fastvideo_args.num_frames > 1 else "image",
|
||||
device=device,
|
||||
extra={},
|
||||
)
|
||||
|
||||
# Run inference
|
||||
start_time = time.time()
|
||||
samples = self.pipeline.forward(
|
||||
batch=batch,
|
||||
fastvideo_args=fastvideo_args,
|
||||
).output
|
||||
|
||||
gen_time = time.time() - start_time
|
||||
logger.info(f"Generated successfully in {gen_time:.2f} seconds")
|
||||
|
||||
# Process outputs
|
||||
videos = rearrange(samples, "b c t h w -> t b c h w")
|
||||
frames = []
|
||||
for x in videos:
|
||||
x = torchvision.utils.make_grid(x, nrow=6)
|
||||
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
|
||||
frames.append((x * 255).numpy().astype(np.uint8))
|
||||
|
||||
# Save video if requested
|
||||
if save_video:
|
||||
save_path = output_path or fastvideo_args.output_path
|
||||
if save_path:
|
||||
os.makedirs(os.path.dirname(save_path), exist_ok=True)
|
||||
video_path = os.path.join(save_path, f"{prompt[:100]}.mp4")
|
||||
imageio.mimsave(video_path, frames, fps=fastvideo_args.fps)
|
||||
logger.info(f"Saved video to {video_path}")
|
||||
else:
|
||||
logger.warning("No output path provided, video not saved")
|
||||
|
||||
if return_frames:
|
||||
return frames
|
||||
else:
|
||||
return {
|
||||
"samples": samples,
|
||||
"prompts": prompt,
|
||||
"size":
|
||||
(target_height, target_width, fastvideo_args.num_frames),
|
||||
"generation_time": gen_time
|
||||
}
|
||||
|
||||
def batch_generate(self,
|
||||
prompts: List[str],
|
||||
output_path: Optional[str] = None,
|
||||
**kwargs) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Generate videos for a batch of prompts.
|
||||
|
||||
Args:
|
||||
prompts: List of prompts to generate videos for
|
||||
output_path: Path to save the videos (overrides the one in fastvideo_args)
|
||||
**kwargs: Additional parameters to pass to generate_video
|
||||
|
||||
Returns:
|
||||
List of output dictionaries from each generation
|
||||
"""
|
||||
if output_path:
|
||||
self.fastvideo_args.output_path = output_path
|
||||
|
||||
results = []
|
||||
for prompt in prompts:
|
||||
result = self.generate_video(prompt=prompt, **kwargs)
|
||||
results.append(result)
|
||||
|
||||
return results
|
||||
|
||||
def image_to_video(self,
|
||||
image: Union[str, torch.Tensor, np.ndarray],
|
||||
prompt: Optional[str] = None,
|
||||
strength: float = 0.8,
|
||||
**kwargs) -> Dict[str, Any]:
|
||||
"""
|
||||
Generate a video from an initial image.
|
||||
|
||||
Args:
|
||||
image: Input image (path, tensor, or numpy array)
|
||||
prompt: Text prompt to guide the generation
|
||||
strength: How much to transform the original image (0-1)
|
||||
**kwargs: Additional parameters to pass to generate_video
|
||||
|
||||
Returns:
|
||||
Output dictionary from the generation
|
||||
"""
|
||||
# Load image if path is provided
|
||||
if isinstance(image, str):
|
||||
# Implementation would depend on your image loading utilities
|
||||
# This is a placeholder for the concept
|
||||
image_tensor = self._load_image(image)
|
||||
elif isinstance(image, np.ndarray):
|
||||
# Convert numpy array to tensor
|
||||
image_tensor = torch.from_numpy(image).permute(2, 0, 1) / 255.0
|
||||
else:
|
||||
image_tensor = image
|
||||
|
||||
# Add image to inference args
|
||||
fastvideo_args = self.fastvideo_args.copy()
|
||||
fastvideo_args.init_image = image_tensor
|
||||
fastvideo_args.strength = strength
|
||||
|
||||
# Generate video
|
||||
return self.generate_video(prompt=prompt or "",
|
||||
fastvideo_args=fastvideo_args,
|
||||
**kwargs)
|
||||
|
||||
def to(self, device: Union[str, torch.device]) -> "VideoGenerator":
|
||||
"""
|
||||
Move the model to the specified device.
|
||||
|
||||
Args:
|
||||
device: The device to move the model to
|
||||
|
||||
Returns:
|
||||
Self for chaining
|
||||
"""
|
||||
device_str = str(device)
|
||||
self.fastvideo_args.device_str = device_str
|
||||
self.fastvideo_args.device = torch.device(device_str)
|
||||
|
||||
# Move pipeline components to device
|
||||
self.pipeline.to(device)
|
||||
|
||||
return self
|
||||
|
||||
def _load_image(self, image_path: str) -> torch.Tensor:
|
||||
"""
|
||||
Load an image from a path and convert to tensor.
|
||||
|
||||
Args:
|
||||
image_path: Path to the image
|
||||
|
||||
Returns:
|
||||
Tensor representation of the image
|
||||
"""
|
||||
# Placeholder implementation - would need to be implemented
|
||||
# based on your image loading utilities
|
||||
import PIL.Image
|
||||
from torchvision import transforms
|
||||
|
||||
image = PIL.Image.open(image_path).convert("RGB")
|
||||
transform = transforms.Compose([
|
||||
transforms.ToTensor(),
|
||||
transforms.Normalize([0.5, 0.5, 0.5], [0.5, 0.5, 0.5])
|
||||
])
|
||||
return transform(image).unsqueeze(0)
|
||||
|
||||
|
||||
def load_prompts_from_file(prompt_path: str) -> List[str]:
|
||||
"""
|
||||
Load prompts from a file.
|
||||
|
||||
Args:
|
||||
prompt_path: Path to the file containing prompts
|
||||
|
||||
Returns:
|
||||
List of prompts
|
||||
"""
|
||||
with open(prompt_path) as f:
|
||||
return [line.strip() for line in f.readlines()]
|
||||
@@ -0,0 +1,20 @@
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
def main():
|
||||
|
||||
# This will automatically handle distributed setup if num_gpus > 1
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/FastHunyuan-Diffusers",
|
||||
num_gpus=2,
|
||||
distributed_executor_backend="mp",
|
||||
)
|
||||
|
||||
# Generate videos with the same simple API, regardless of GPU count
|
||||
prompt = "A beautiful woman in a red dress walking down a street"
|
||||
video = generator.generate_video(prompt)
|
||||
|
||||
prompt2 = "A beautiful woman in a blue dress walking down a street"
|
||||
video2 = generator.generate_video(prompt2)
|
||||
|
||||
if __name__ == '__main__':
|
||||
main()
|
||||
@@ -19,7 +19,7 @@ class FastVideoArgs:
|
||||
model_path: str
|
||||
|
||||
# Distributed executor backend
|
||||
distributed_executor_backend: str = "torch"
|
||||
distributed_executor_backend: str = "mp"
|
||||
|
||||
inference_mode: bool = True # if False == training mode
|
||||
|
||||
@@ -121,12 +121,19 @@ class FastVideoArgs:
|
||||
type=str,
|
||||
help="Directory containing StepVideo model",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--distributed-executor-backend",
|
||||
type=str,
|
||||
default=FastVideoArgs.distributed_executor_backend,
|
||||
choices=["torch"],
|
||||
help="Backend for distributed execution",
|
||||
)
|
||||
|
||||
# distributed_executor_backend
|
||||
parser.add_argument(
|
||||
"--distributed-executor-backend",
|
||||
type=str,
|
||||
choices=["mp", "ray", "torch"],
|
||||
choices=["mp", "torch"],
|
||||
default=FastVideoArgs.distributed_executor_backend,
|
||||
help="The distributed executor backend to use",
|
||||
)
|
||||
@@ -431,7 +438,7 @@ class FastVideoArgs:
|
||||
|
||||
return cls(**kwargs)
|
||||
|
||||
def check_inference_args(self) -> None:
|
||||
def check_fastvideo_args(self) -> None:
|
||||
"""Validate inference arguments for consistency"""
|
||||
if self.tp_size is None:
|
||||
self.tp_size = self.num_gpus
|
||||
@@ -470,7 +477,7 @@ def prepare_fastvideo_args(argv: List[str]) -> FastVideoArgs:
|
||||
FastVideoArgs.add_cli_args(parser)
|
||||
raw_args = parser.parse_args(argv)
|
||||
fastvideo_args = FastVideoArgs.from_cli_args(raw_args)
|
||||
fastvideo_args.check_inference_args()
|
||||
fastvideo_args.check_fastvideo_args()
|
||||
global _current_fastvideo_args
|
||||
_current_fastvideo_args = fastvideo_args
|
||||
return fastvideo_args
|
||||
|
||||
+69
-1
@@ -4,16 +4,21 @@
|
||||
import argparse
|
||||
import hashlib
|
||||
import importlib
|
||||
import ctypes
|
||||
import signal
|
||||
import inspect
|
||||
import json
|
||||
import math
|
||||
import traceback
|
||||
import os
|
||||
import sys
|
||||
import tempfile
|
||||
from functools import wraps, partial
|
||||
from typing import Any, Dict, List, Optional, Type, TypeVar, Union, cast, Callable
|
||||
from typing import Any, Dict, List, Optional, Type, TypeVar, Union, cast, Callable, Tuple
|
||||
from dataclasses import asdict, fields
|
||||
import cloudpickle
|
||||
import zmq
|
||||
import psutil
|
||||
|
||||
import filelock
|
||||
import torch
|
||||
@@ -554,3 +559,66 @@ def update_in_place(target, source, ignore_fields=()):
|
||||
for f in fields(target):
|
||||
if hasattr(source, f.name) and f.name not in list(ignore_fields):
|
||||
setattr(target, f.name, getattr(source, f.name))
|
||||
|
||||
|
||||
def get_zmq_socket(
|
||||
context: zmq.Context, socket_type: zmq.SocketType, endpoint: str, bind: bool
|
||||
):
|
||||
mem = psutil.virtual_memory()
|
||||
total_mem = mem.total / 1024**3
|
||||
available_mem = mem.available / 1024**3
|
||||
if total_mem > 32 and available_mem > 16:
|
||||
buf_size = int(0.5 * 1024**3)
|
||||
else:
|
||||
buf_size = -1
|
||||
|
||||
socket = context.socket(socket_type)
|
||||
|
||||
def set_send_opt():
|
||||
socket.setsockopt(zmq.SNDHWM, 0)
|
||||
socket.setsockopt(zmq.SNDBUF, buf_size)
|
||||
|
||||
def set_recv_opt():
|
||||
socket.setsockopt(zmq.RCVHWM, 0)
|
||||
socket.setsockopt(zmq.RCVBUF, buf_size)
|
||||
|
||||
if socket_type == zmq.PUSH:
|
||||
set_send_opt()
|
||||
elif socket_type == zmq.PULL:
|
||||
set_recv_opt()
|
||||
elif socket_type == zmq.DEALER:
|
||||
set_send_opt()
|
||||
set_recv_opt()
|
||||
else:
|
||||
raise ValueError(f"Unsupported socket type: {socket_type}")
|
||||
|
||||
if bind:
|
||||
socket.bind(endpoint)
|
||||
else:
|
||||
socket.connect(endpoint)
|
||||
|
||||
return socket
|
||||
|
||||
def kill_itself_when_parent_died():
|
||||
# if sys.platform == "linux":
|
||||
# sigkill this process when parent worker manager dies
|
||||
PR_SET_PDEATHSIG = 1
|
||||
libc = ctypes.CDLL("libc.so.6")
|
||||
libc.prctl(PR_SET_PDEATHSIG, signal.SIGKILL)
|
||||
# else:
|
||||
# logger.warning("kill_itself_when_parent_died is only supported in linux.")
|
||||
|
||||
def get_exception_traceback():
|
||||
etype, value, tb = sys.exc_info()
|
||||
err_str = "".join(traceback.format_exception(etype, value, tb))
|
||||
return err_str
|
||||
|
||||
class TypeBasedDispatcher:
|
||||
def __init__(self, mapping: List[Tuple[Type, Callable]]):
|
||||
self._mapping = mapping
|
||||
|
||||
def __call__(self, obj: Any):
|
||||
for ty, fn in self._mapping:
|
||||
if isinstance(obj, ty):
|
||||
return fn(obj)
|
||||
raise ValueError(f"Invalid object: {obj}")
|
||||
@@ -0,0 +1,74 @@
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Callable, List, Tuple, Any, Optional, Union, TypeVar, Dict
|
||||
|
||||
import torch
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.pipelines import ForwardBatch
|
||||
from fastvideo.v1.utils import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
_R = TypeVar("_R")
|
||||
|
||||
class Executor(ABC):
|
||||
|
||||
def __init__(self, fastvideo_args: FastVideoArgs):
|
||||
self.fastvideo_args = fastvideo_args
|
||||
|
||||
self._init_executor()
|
||||
|
||||
@abstractmethod
|
||||
def _init_executor(self) -> None:
|
||||
raise NotImplementedError
|
||||
|
||||
@classmethod
|
||||
def get_class(cls, fastvideo_args: FastVideoArgs) -> type["Executor"]:
|
||||
if fastvideo_args.distributed_executor_backend == "torch":
|
||||
raise ValueError("Torch is not supported yet")
|
||||
from fastvideo.v1.worker.torchrun_executor import TorchRunExecutor
|
||||
return TorchRunExecutor
|
||||
elif fastvideo_args.distributed_executor_backend == "mp":
|
||||
from fastvideo.v1.worker.multiproc_executor import MultiprocExecutor
|
||||
return MultiprocExecutor
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Unsupported distributed executor backend: {fastvideo_args.distributed_executor_backend}"
|
||||
)
|
||||
|
||||
def execute_model(
|
||||
self,
|
||||
forward_batch: ForwardBatch,
|
||||
) -> ForwardBatch:
|
||||
output = self.collective_rpc("execute_forward",
|
||||
args=(forward_batch, ))
|
||||
return output
|
||||
|
||||
@abstractmethod
|
||||
def collective_rpc(self,
|
||||
method: Union[str, Callable[..., _R]],
|
||||
timeout: Optional[float] = None,
|
||||
args: Tuple = (),
|
||||
kwargs: Optional[Dict[str, Any]] = None) -> List[_R]:
|
||||
"""
|
||||
Execute an RPC call on all workers.
|
||||
|
||||
Args:
|
||||
method: Name of the worker method to execute, or a callable that
|
||||
is serialized and sent to all workers to execute.
|
||||
|
||||
If the method is a callable, it should accept an additional
|
||||
`self` argument, in addition to the arguments passed in `args`
|
||||
and `kwargs`. The `self` argument will be the worker object.
|
||||
timeout: Maximum time in seconds to wait for execution. Raises a
|
||||
:exc:`TimeoutError` on timeout. `None` means wait indefinitely.
|
||||
args: Positional arguments to pass to the worker method.
|
||||
kwargs: Keyword arguments to pass to the worker method.
|
||||
|
||||
Returns:
|
||||
A list containing the results from each worker.
|
||||
|
||||
Note:
|
||||
It is recommended to use this API to only pass control messages,
|
||||
and set up data-plane communication to pass data.
|
||||
"""
|
||||
raise NotImplementedError
|
||||
@@ -0,0 +1,230 @@
|
||||
from abc import ABC
|
||||
|
||||
import torch
|
||||
|
||||
from typing import Optional, List
|
||||
import multiprocessing as mp
|
||||
import setproctitle
|
||||
import psutil
|
||||
import faulthandler
|
||||
import signal
|
||||
import traceback
|
||||
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs, set_current_fastvideo_args
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines import ForwardBatch
|
||||
from fastvideo.v1.utils import run_method, update_environment_variables
|
||||
from fastvideo.v1.distributed import (
|
||||
init_distributed_environment,
|
||||
initialize_model_parallel,
|
||||
)
|
||||
import zmq
|
||||
from fastvideo.v1.utils import get_zmq_socket, kill_itself_when_parent_died, get_exception_traceback, TypeBasedDispatcher
|
||||
from fastvideo.v1.worker.io_struct import RpcReqInput, RpcReqOutput, GenerateRequest
|
||||
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
# class WorkerWrapper5
|
||||
# def __init__(self, fastvideo_args: FastVideoArgs):
|
||||
# self.fastvideo_args = fastvideo_args
|
||||
|
||||
# def update_environment_variables(self, envs_list: List[Dict[str, str]]) -> None:
|
||||
# envs = envs_list[self.rpc_rank]
|
||||
# key = 'CUDA_VISIBLE_DEVICES'
|
||||
# if key in envs and key in os.environ:
|
||||
# # overwriting CUDA_VISIBLE_DEVICES is desired behavior
|
||||
# # suppress the warning in `update_environment_variables`
|
||||
# del os.environ[key]
|
||||
# update_environment_variables(envs)
|
||||
|
||||
# def init_worker(self, all_kwargs: List[Dict[str, Any]]) -> None:
|
||||
# kwargs = all_kwargs[self.rpc_rank]
|
||||
# self.fastvideo_args = kwargs.get("fastvideo_args", None)
|
||||
# assert self.fastvideo_args is not None, (
|
||||
# "fastvideo_args is required to initialize the worker")
|
||||
|
||||
# with set_current_fastvideo_args(self.fastvideo_args):
|
||||
# self.worker = Worker(**kwargs)
|
||||
# assert self.worker is not None
|
||||
|
||||
# def init_device(self):
|
||||
# self.worker.init_device()
|
||||
|
||||
# def execute_method(self, method: Union[str, bytes], *args, **kwargs):
|
||||
# try:
|
||||
# # method resolution order:
|
||||
# # if a method is defined in this class, it will be called directly.
|
||||
# # otherwise, since we define `__getattr__` and redirect attribute
|
||||
# # query to `self.worker`, the method will be called on the worker.
|
||||
# return run_method(self, method, args, kwargs)
|
||||
# except Exception as e:
|
||||
# # if the driver worker also execute methods,
|
||||
# # exceptions in the rest worker may cause deadlock in rpc like ray
|
||||
# # see https://github.com/vllm-project/vllm/issues/3455
|
||||
# # print the error and inform the user to solve the error
|
||||
# msg = (f"Error executing method {method!r}. "
|
||||
# "This might cause deadlock in distributed execution.")
|
||||
# logger.exception(msg)
|
||||
# raise e
|
||||
|
||||
# def __getattr__(self, attr):
|
||||
# return getattr(self.worker, attr)
|
||||
|
||||
|
||||
class Worker(ABC):
|
||||
|
||||
def __init__(self, fastvideo_args: FastVideoArgs, local_rank: int,
|
||||
rank: int, is_driver_worker: bool = False):
|
||||
self.fastvideo_args = fastvideo_args
|
||||
self.local_rank = local_rank
|
||||
self.rank = rank
|
||||
# TODO: don't hardcode this
|
||||
self.distributed_init_method = "env://"
|
||||
self.is_driver_worker = is_driver_worker
|
||||
|
||||
context = zmq.Context(2)
|
||||
self.recv_from_rpc = get_zmq_socket(context, zmq.DEALER, "ipc://fastvideo_rpc_broadcast", False)
|
||||
# self.send_to_rpc = get_zmq_socket(context, zmq.PUSH, "inproc://fastvideo_rpc_broadcast", True)
|
||||
|
||||
# Init request dispatcher
|
||||
self._request_dispatcher = TypeBasedDispatcher(
|
||||
[
|
||||
# (TokenizedGenerateReqInput, self.handle_generate_request),
|
||||
# (TokenizedEmbeddingReqInput, self.handle_embedding_request),
|
||||
# (FlushCacheReq, self.flush_cache_wrapped),
|
||||
# (AbortReq, self.abort_request),
|
||||
# (OpenSessionReqInput, self.open_session),
|
||||
# (CloseSessionReqInput, self.close_session),
|
||||
# (UpdateWeightFromDiskReqInput, self.update_weights_from_disk),
|
||||
# (InitWeightsUpdateGroupReqInput, self.init_weights_update_group),
|
||||
# (
|
||||
# UpdateWeightsFromDistributedReqInput,
|
||||
# self.update_weights_from_distributed,
|
||||
# ),
|
||||
# (UpdateWeightsFromTensorReqInput, self.update_weights_from_tensor),
|
||||
# (GetWeightsByNameReqInput, self.get_weights_by_name),
|
||||
# (ReleaseMemoryOccupationReqInput, self.release_memory_occupation),
|
||||
# (ResumeMemoryOccupationReqInput, self.resume_memory_occupation),
|
||||
# (ProfileReq, self.profile),
|
||||
# (GetInternalStateReq, self.get_internal_state),
|
||||
# (SetInternalStateReq, self.set_internal_state),
|
||||
(RpcReqInput, self.handle_rpc_request),
|
||||
(GenerateRequest, self.handle_generate_request),
|
||||
# (ExpertDistributionReq, self.expert_distribution_handle),
|
||||
]
|
||||
)
|
||||
|
||||
def handle_rpc_request(self, req: RpcReqInput) -> RpcReqOutput:
|
||||
pass
|
||||
|
||||
def handle_generate_request(self, req: GenerateRequest) -> None:
|
||||
logger.info(f"Worker {self.rank} received generate request")
|
||||
pass
|
||||
|
||||
def init_device(self):
|
||||
if self.fastvideo_args.device_str.startswith("cuda"):
|
||||
# torch.distributed.all_reduce does not free the input tensor until
|
||||
# the synchronization point. This causes the memory usage to grow
|
||||
# as the number of all_reduce calls increases. This env var disables
|
||||
# this behavior.
|
||||
# Related issue:
|
||||
# https://discuss.pytorch.org/t/cuda-allocation-lifetime-for-inputs-to-distributed-all-reduce/191573
|
||||
os.environ["TORCH_NCCL_AVOID_RECORD_STREAMS"] = "1"
|
||||
|
||||
# This env var set by Ray causes exceptions with graph building.
|
||||
os.environ.pop("NCCL_ASYNC_ERROR_HANDLING", None)
|
||||
self.device = torch.device(f"cuda:{self.local_rank}")
|
||||
torch.cuda.set_device(self.device)
|
||||
|
||||
_check_if_gpu_supports_dtype(self.model_config.dtype)
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
self.init_gpu_memory = torch.cuda.mem_get_info()[0]
|
||||
else:
|
||||
raise ValueError(f"Unsupported device: {self.fastvideo_args.device_str}")
|
||||
|
||||
# Initialize the distributed environment.
|
||||
init_worker_distributed_environment(self.fastvideo_args, self.rank,
|
||||
self.distributed_init_method,
|
||||
self.local_rank)
|
||||
|
||||
self.pipeline = build_pipeline(self.fastvideo_args)
|
||||
|
||||
def execute_forward(self, forward_batch: ForwardBatch) -> torch.Tensor:
|
||||
return self.pipeline.forward(forward_batch)
|
||||
|
||||
def recv_requests(self) -> List[ForwardBatch]:
|
||||
pass
|
||||
|
||||
def send_response(self, output: torch.Tensor) -> None:
|
||||
pass
|
||||
|
||||
def event_loop(self) -> None:
|
||||
"""Event loop for the worker."""
|
||||
logger.info(f"Worker {self.rank} starting event loop...")
|
||||
recv_rpc = self.recv_from_rpc.recv_pyobj()
|
||||
assert recv_rpc == "wait_until_ready"
|
||||
logger.info(f"Worker {self.rank} received wait_until_ready")
|
||||
self.recv_from_rpc.send_pyobj(f"ready{self.rank}")
|
||||
logger.info(f"Worker {self.rank} sent ready")
|
||||
logger.info(f"Worker {self.rank} started event loop")
|
||||
while True:
|
||||
logger.info(f"Worker {self.rank} waiting for RPC")
|
||||
recv_rpc = self.recv_from_rpc.recv_pyobj()
|
||||
logger.info(f"Received RPC: {recv_rpc}")
|
||||
# assert isinstance(recv_rpc, RpcReqInput)
|
||||
inputs = self._request_dispatcher(recv_rpc)
|
||||
output = self.execute_forward(inputs)
|
||||
self.send_response(output)
|
||||
|
||||
|
||||
def init_worker_distributed_environment(
|
||||
fastvideo_args: FastVideoArgs,
|
||||
rank: int,
|
||||
distributed_init_method: Optional[str] = None,
|
||||
local_rank: int = -1,
|
||||
) -> None:
|
||||
"""Initialize distributed environment and model parallelism."""
|
||||
|
||||
world_size = fastvideo_args.num_gpus
|
||||
|
||||
torch.cuda.set_device(local_rank)
|
||||
init_distributed_environment(world_size=world_size,
|
||||
rank=rank,
|
||||
local_rank=local_rank)
|
||||
device_str = f"cuda:{local_rank}"
|
||||
fastvideo_args.device_str = device_str
|
||||
fastvideo_args.device = torch.device(device_str)
|
||||
assert fastvideo_args.sp_size is not None
|
||||
assert fastvideo_args.tp_size is not None
|
||||
initialize_model_parallel(
|
||||
sequence_model_parallel_size=fastvideo_args.sp_size,
|
||||
tensor_model_parallel_size=fastvideo_args.tp_size,
|
||||
)
|
||||
|
||||
|
||||
def run_worker_process(fastvideo_args: FastVideoArgs, local_rank: int,
|
||||
rank: int, pipe_writer: mp.Pipe):
|
||||
print(f"Worker {rank} starting...")
|
||||
try:
|
||||
# Config the process
|
||||
kill_itself_when_parent_died()
|
||||
prefix = f"{local_rank}"
|
||||
setproctitle.setproctitle(f"fastvideo::gpu_worker{prefix.replace(' ', '_')}")
|
||||
faulthandler.enable()
|
||||
parent_process = psutil.Process().parent()
|
||||
|
||||
logger.info(f"Worker {rank} initializing...")
|
||||
worker = Worker(fastvideo_args, local_rank, rank, False)
|
||||
pipe_writer.send(
|
||||
{
|
||||
"status": "ready",
|
||||
"local_rank": local_rank,
|
||||
}
|
||||
)
|
||||
worker.event_loop()
|
||||
except Exception:
|
||||
traceback = get_exception_traceback()
|
||||
logger.error(f"Worker {rank} hit an exception: {traceback}")
|
||||
parent_process.send_signal(signal.SIGQUIT)
|
||||
@@ -0,0 +1,18 @@
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional, Dict
|
||||
|
||||
@dataclass
|
||||
class RpcReqInput:
|
||||
method: str
|
||||
parameters: Optional[Dict] = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class RpcReqOutput:
|
||||
success: bool
|
||||
message: str
|
||||
|
||||
@dataclass
|
||||
class GenerateRequest:
|
||||
prompt: str
|
||||
seed: int
|
||||
@@ -0,0 +1,115 @@
|
||||
import signal
|
||||
import psutil
|
||||
import multiprocessing as mp
|
||||
import time
|
||||
import pickle
|
||||
import cloudpickle
|
||||
import zmq
|
||||
from typing import List, Callable, Any, Optional, Union
|
||||
from multiprocessing.process import BaseProcess
|
||||
|
||||
from fastvideo.v1.worker.executor import Executor
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.worker.gpu_worker import run_worker_process
|
||||
from fastvideo.v1.utils import get_zmq_socket
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class MultiprocExecutor(Executor):
|
||||
|
||||
def _init_executor(self) -> None:
|
||||
# The child processes will send SIGUSR1 when unrecoverable
|
||||
# errors happen.
|
||||
def sigusr1_handler(signum, frame):
|
||||
logger.fatal(
|
||||
"MulitprocExecutor got fatal signal from worker processes, "
|
||||
"shutting down. See stack trace above for root cause issue.")
|
||||
# Propagate error up to parent process.
|
||||
parent_process = psutil.Process().parent()
|
||||
parent_process.send_signal(signal.SIGUSR1)
|
||||
self.shutdown()
|
||||
|
||||
signal.signal(signal.SIGUSR1, sigusr1_handler)
|
||||
|
||||
self.world_size = self.fastvideo_args.num_gpus
|
||||
|
||||
# Set mp start method
|
||||
mp.set_start_method("spawn", force=True)
|
||||
|
||||
self.workers: List[BaseProcess] = []
|
||||
self.worker_pipe_readers: List[mp.Pipe] = []
|
||||
for rank in range(self.world_size):
|
||||
reader, writer = mp.Pipe(duplex=False)
|
||||
self.worker_pipe_readers.append(reader)
|
||||
|
||||
worker = mp.Process(target=run_worker_process,
|
||||
args=(self.fastvideo_args, rank, rank,
|
||||
writer))
|
||||
worker.start()
|
||||
self.workers.append(worker)
|
||||
logger.info(f"Workers: {self.workers}")
|
||||
for reader in self.worker_pipe_readers:
|
||||
data = reader.recv()
|
||||
assert data["status"] == "ready"
|
||||
logger.info(f"Worker {data['local_rank']} ready")
|
||||
|
||||
context = zmq.Context(2)
|
||||
self.send_to_rpc = get_zmq_socket(
|
||||
context, zmq.DEALER, "ipc://fastvideo_rpc_broadcast", True
|
||||
)
|
||||
logger.info("Sending wait_until_ready to workers")
|
||||
self.send_to_rpc.send_pyobj("wait_until_ready")
|
||||
self.send_to_rpc.send_pyobj("wait_until_ready")
|
||||
logger.info("Sent wait_until_ready to workers")
|
||||
recv_req = self.send_to_rpc.recv_pyobj(zmq.BLOCKY)
|
||||
print(f"Received wait_until_ready from workers: {recv_req}")
|
||||
recv_req = self.send_to_rpc.recv_pyobj(zmq.BLOCKY)
|
||||
print(f"Received wait_until_ready from workers: {recv_req}")
|
||||
|
||||
|
||||
|
||||
|
||||
# self.collective_rpc("wait_until_ready")
|
||||
|
||||
def collective_rpc(self,
|
||||
method: Union[str, Callable],
|
||||
timeout: Optional[float] = None,
|
||||
args: tuple = (),
|
||||
kwargs: Optional[dict] = None) -> list[Any]:
|
||||
start_time = time.monotonic()
|
||||
kwargs = kwargs or {}
|
||||
|
||||
# NOTE: If the args are heterogeneous, then we pack them into a list,
|
||||
# and unpack them in the method of every worker, because every worker
|
||||
# knows their own rank.
|
||||
try:
|
||||
if isinstance(method, str):
|
||||
send_method = method
|
||||
else:
|
||||
send_method = cloudpickle.dumps(
|
||||
method, protocol=pickle.HIGHEST_PROTOCOL)
|
||||
self.rpc_broadcast_mq.enqueue((send_method, args, kwargs))
|
||||
|
||||
responses = [None] * self.world_size
|
||||
for w in self.workers:
|
||||
dequeue_timeout = timeout - (time.monotonic() - start_time
|
||||
) if timeout is not None else None
|
||||
status, result = w.worker_response_mq.dequeue(
|
||||
timeout=dequeue_timeout)
|
||||
|
||||
if status != WorkerProc.ResponseStatus.SUCCESS:
|
||||
if isinstance(result, Exception):
|
||||
raise result
|
||||
else:
|
||||
raise RuntimeError("Worker failed")
|
||||
|
||||
responses[w.rank] = result
|
||||
|
||||
return responses
|
||||
except TimeoutError as e:
|
||||
raise TimeoutError(f"RPC call to {method} timed out.") from e
|
||||
except Exception as e:
|
||||
# Re-raise any other exceptions
|
||||
raise e
|
||||
@@ -0,0 +1,7 @@
|
||||
class WorkerWrapper:
|
||||
|
||||
def __init__(self, worker: Worker):
|
||||
self.worker = worker
|
||||
|
||||
def __call__(self, *args, **kwargs):
|
||||
return self.worker(*args, **kwargs)
|
||||
Reference in New Issue
Block a user