Compare commits

...
3 Commits
Author SHA1 Message Date
Will Lin 99dff854ad zmq worker 2025-04-16 11:08:29 -07:00
Will Lin 4ccf68d71e checkpoint 2025-04-15 16:05:04 -07:00
Will Lin 15b798a9d1 checkpoint 2025-04-15 16:02:45 -07:00
14 changed files with 965 additions and 12 deletions
+2
View File
@@ -1 +1,3 @@
from fastvideo.v1.entrypoints.video_generator import VideoGenerator
__all__ = ["VideoGenerator"]
+2 -1
View File
@@ -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"
]
-3
View File
@@ -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):
+2 -2
View File
@@ -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):
+13 -1
View File
@@ -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",
]
+402
View File
@@ -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()]
+20
View File
@@ -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()
+11 -4
View File
@@ -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
View File
@@ -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}")
+74
View File
@@ -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
+230
View File
@@ -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)
+18
View File
@@ -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
+115
View File
@@ -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
+7
View File
@@ -0,0 +1,7 @@
class WorkerWrapper:
def __init__(self, worker: Worker):
self.worker = worker
def __call__(self, *args, **kwargs):
return self.worker(*args, **kwargs)