Not a member of Pastebin yet?
Sign Up,
it unlocks many cool features!
- import os
- import time
- import torch
- import torchvision
- from torch import Tensor
- from transformers import AutoTokenizer, T5Gemma2Model, BatchEncoding
- from diffusers import AutoencoderKLWan, FlowMatchEulerDiscreteScheduler
- from transformer.transformer_motif_video import MotifVideoTransformer3DModel
- # Parameters
- prompt = "A category-five hurricane, viewed from inside the eye, reveals a circular stadium of cloud walls rising to fifty thousand feet with an eerie disk of blue sky directly overhead. Shot from a NOAA reconnaissance aircraft mounted camera, the perspective looks outward toward the eyewall — a near-vertical curtain of rotating cloud and lightning that is simultaneously terrifying and transcendent. The inner surface of the eyewall catches the setting sun, painting it in improbable shades of peach and rose. The camera slowly pans 360 degrees to complete one full revolution, capturing the entire coliseum of the storm. Below, the ocean surface is a white blur of foam and spray. The documentary-style cinematography strips away all artifice to present the storm as an entity of pure elemental power."
- negative_prompt = ""
- height = 736
- width = 1280
- num_frames = 121
- num_inference_steps = 50
- #cfg = ?
- device = "cuda"
- dtype = torch.bfloat16
- model_path = "./" # Ensure this points to the root of your local Motif-Video-2B repo
- # Load Components
- tokenizer = AutoTokenizer.from_pretrained(model_path, subfolder="tokenizer")
- text_encoder = T5Gemma2Model.from_pretrained(model_path, subfolder="text_encoder", torch_dtype=dtype)
- vae = AutoencoderKLWan.from_pretrained(model_path, subfolder="vae", torch_dtype=dtype)
- transformer = MotifVideoTransformer3DModel.from_pretrained(model_path, subfolder="transformer", torch_dtype=dtype)
- scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained(model_path, subfolder="scheduler")
- text_encoder.eval()
- vae.eval()
- transformer.eval()
- # ==========================================
- # Util
- # ==========================================
- # Copied from diffusers.pipelines.flux.pipeline_flux.calculate_shift
- '''def calculate_shift(image_seq_len, base_seq_len: int = 256, max_seq_len: int = 4096, base_shift: float = 0.5, max_shift: float = 1.15):
- m = (max_shift - base_shift) / (max_seq_len - base_seq_len)
- b = base_shift - m * base_seq_len
- mu = image_seq_len * m + b
- return mu
- # Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion.retrieve_timesteps
- def retrieve_timesteps(scheduler, num_inference_steps: Optional[int] = None, device: Optional[Union[str, torch.device]] = None, timesteps: Optional[List[int]] = None, sigmas: Optional[List[float]] = None, use_linear_quadratic_schedule: bool = False, linear_quadratic_emulating_steps: int = 250, **kwargs):
- """
- Retrieve timesteps from the scheduler.
- Args:
- scheduler: The noise scheduler to use.
- num_inference_steps: Number of denoising steps.
- device: Device to place timesteps on.
- timesteps: Custom timestep values (mutually exclusive with sigmas).
- sigmas: Custom sigma values (mutually exclusive with timesteps).
- use_linear_quadratic_schedule: If True, use linear-quadratic sigma schedule.
- This overrides the default linear schedule. Requires num_inference_steps
- to be even.
- linear_quadratic_emulating_steps: Controls the linear portion slope.
- Higher values result in gentler slope in the first half. Default: 250.
- **kwargs: Additional arguments passed to scheduler.set_timesteps().
- Returns:
- Tuple of (timesteps, num_inference_steps).
- Raises:
- ValueError: If both timesteps and sigmas are provided, or if
- use_linear_quadratic_schedule is True but num_inference_steps is odd.
- """
- if timesteps is not None and sigmas is not None:
- raise ValueError("Only one of `timesteps` or `sigmas` can be passed. Please choose one to set custom values")
- # Handle linear-quadratic schedule: compute sigmas if flag is set
- if use_linear_quadratic_schedule:
- if sigmas is not None:
- raise ValueError(
- "Cannot use both `sigmas` and `use_linear_quadratic_schedule`. "
- "The linear-quadratic schedule computes sigmas automatically."
- )
- if num_inference_steps is None:
- raise ValueError("`num_inference_steps` must be provided when using `use_linear_quadratic_schedule`.")
- sigmas = get_linear_quadratic_sigmas(
- num_inference_steps=num_inference_steps,
- linear_quadratic_emulating_steps=linear_quadratic_emulating_steps,
- )
- if timesteps is not None:
- accepts_timesteps = "timesteps" in set(inspect.signature(scheduler.set_timesteps).parameters.keys())
- if not accepts_timesteps:
- raise ValueError(
- f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom"
- f" timestep schedules. Please check whether you are using the correct scheduler."
- )
- scheduler.set_timesteps(timesteps=timesteps, device=device, **kwargs)
- timesteps = scheduler.timesteps
- num_inference_steps = len(timesteps)
- elif sigmas is not None:
- accept_sigmas = "sigmas" in set(inspect.signature(scheduler.set_timesteps).parameters.keys())
- if not accept_sigmas:
- raise ValueError(
- f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom"
- f" sigmas schedules. Please check whether you are using the correct scheduler."
- )
- scheduler.set_timesteps(sigmas=sigmas, device=device, **kwargs)
- timesteps = scheduler.timesteps
- num_inference_steps = len(timesteps)
- else:
- scheduler.set_timesteps(num_inference_steps, device=device, **kwargs)
- timesteps = scheduler.timesteps
- return timesteps, num_inference_steps
- def get_linear_quadratic_sigmas(num_inference_steps: int, linear_quadratic_emulating_steps: int = 250) -> np.ndarray:
- """
- Compute a linear-quadratic sigma schedule for flow matching.
- This schedule combines:
- - First half: Linear interpolation from high noise to medium noise (slow denoising)
- - Second half: Quadratic interpolation from medium noise to clean (faster denoising)
- Convention:
- - sigma=1.0 represents pure noise
- - sigma=0.0 represents clean image
- - Output sigmas are in descending order (1.0 → ~0)
- Args:
- num_inference_steps: Total number of denoising steps (must be even).
- linear_quadratic_emulating_steps: Controls the slope of linear interpolation.
- Higher values result in gentler slope in the first half.
- Returns:
- np.ndarray: Array of sigma values with shape (num_inference_steps,).
- The scheduler will append a terminal 0.
- Raises:
- ValueError: If num_inference_steps is not even.
- Reference:
- Linear-quadratic timestep schedule for improved flow matching inference.
- """
- if num_inference_steps % 2 != 0:
- raise ValueError(
- f"num_inference_steps must be even for linear-quadratic schedule, but got {num_inference_steps}"
- )
- steps = num_inference_steps
- N = linear_quadratic_emulating_steps
- half_steps = steps // 2
- # First half: linear interpolation from 1 toward 0
- # Takes first half_steps values from linspace(1, 0, N+1)
- linear_part = np.linspace(1.0, 0.0, N + 1)[:half_steps]
- # Second half: quadratic interpolation
- # Formula: x^2 * (half_steps/N - 1) - (half_steps/N - 1)
- # = (half_steps/N - 1) * (x^2 - 1)
- # This maps x=0 to (half_steps/N - 1) * (-1) = 1 - half_steps/N
- # and maps x=1 to 0
- x = np.linspace(0.0, 1.0, half_steps + 1)
- scale_factor = half_steps / N - 1 # negative value
- quadratic_part = x**2 * scale_factor - scale_factor
- # Concatenate and exclude the last 0 (scheduler appends terminal 0)
- sigmas = np.concatenate([linear_part, quadratic_part])
- sigmas = sigmas[:-1] # Remove trailing 0, scheduler will append it
- return sigmas.astype(np.float32)
- '''
- # ==========================================
- # Step 1: Encode Text
- # ==========================================
- def encode_prompt(prompt, tokenizer, text_encoder, device):
- prompt = [prompt] if isinstance(prompt, str) else prompt
- text_inputs = tokenizer(
- prompt,
- padding="max_length",
- max_length=512,
- truncation=True,
- add_special_tokens=True,
- return_attention_mask=True,
- return_tensors="pt",
- )
- #input_ids = text_inputs.input_ids.to(device)
- #attention_mask = text_inputs.attention_mask.to(device)
- text_inputs = BatchEncoding(
- {k: v.to(device) if isinstance(v, torch.Tensor) else v for k, v in text_inputs.items()}
- )
- attention_mask = text_inputs.attention_mask.to(device)
- '''# Encode text sequence
- # T5Gemma2Model bundles encoder and decoder/LM head, while _get_default_embeds expects an encoder-only model (similar to T5EncoderModel/T5GemmaEncoderModel), so use the encoder submodule explicitly here
- outputs = text_encoder.encoder(input_ids=input_ids, attention_mask=attention_mask, output_hidden_states=False, return_dict=True)
- # T5Gemma2Model returns an encoder_last_hidden_state or last_hidden_state depending on version — check both:
- if hasattr(outputs, "encoder_last_hidden_state") and outputs.encoder_last_hidden_state is not None:
- prompt_embeds = outputs.encoder_last_hidden_state
- else:
- prompt_embeds = outputs.last_hidden_state'''
- prompt_embeds = text_encoder.encoder(**text_inputs)[0]
- prompt_embeds = prompt_embeds.to(device=device)
- pooled_prompt_embeds = average_pool(prompt_embeds, attention_mask)
- return prompt_embeds, attention_mask, pooled_prompt_embeds
- def average_pool(last_hidden_states: Tensor, attention_mask: Tensor) -> Tensor:
- last_hidden = last_hidden_states.masked_fill(~attention_mask[..., None].bool(), 0.0)
- denom = attention_mask.sum(dim=1, keepdim=True).clamp(min=1) # avoid div by zero
- return last_hidden.sum(dim=1) / denom
- print("Encoding prompt...")
- text_encoder.to(device)
- with torch.no_grad():
- prompt_embeds, attn_mask, pooled_prompt_embeds = encode_prompt(prompt, tokenizer, text_encoder, device)
- neg_prompt_embeds, neg_attn_mask, neg_pooled_prompt_embeds = encode_prompt(negative_prompt, tokenizer, text_encoder, device)
- encoder_hidden_states = torch.cat([neg_prompt_embeds, prompt_embeds], dim=0).to(device, dtype)
- encoder_attention_mask = torch.cat([neg_attn_mask, attn_mask], dim=0).to(device, dtype)
- pooled_projections = torch.cat([neg_pooled_prompt_embeds, pooled_prompt_embeds], dim=0).to(device, dtype)
- text_encoder.to("cpu")
- torch.cuda.empty_cache()
- # ==========================================
- # Step 2: Denoising Loop
- # ==========================================
- transformer.to(device)
- print("Preparing latents...")
- # Motif-Video-2B utilizes the Wan 3D VAE structure:
- # Temporal downsampling factor is 4, Spatial downsampling factor is 8.
- latent_frames = (num_frames - 1) // 4 + 1
- latent_height = height // 8
- latent_width = width // 8
- # VAE z_dim is typically 16 channels
- #channels = getattr(transformer.config, "in_channels", 16)
- channels = getattr(transformer.config, "in_channels", vae.config.z_dim) # can read directly?
- shape = (1, channels, latent_frames, latent_height, latent_width)
- latents = torch.randn(shape, generator=torch.Generator(device=device), device=device, dtype=dtype)
- print("Preparing scheduler...")
- scheduler.set_timesteps(num_inference_steps, device=device)
- timesteps = scheduler.timesteps
- print("Starting denoising loop...")
- for i, t in enumerate(timesteps):
- print(f"Step {i+1}/{num_inference_steps} (timestep {t.item()})")
- with torch.no_grad():
- noise_pred = transformer(
- hidden_states=latents,
- timestep=t.unsqueeze(0).expand(latents.shape[0]), # Expand scalar timestep to 1D tensor
- attention_kwargs=None,
- return_dict=False,
- tread_disabled=True,
- encoder_hidden_states=encoder_hidden_states,
- encoder_attention_mask=encoder_attention_mask,
- pooled_projections=pooled_projections,
- image_embeds=None,
- )[0]
- # Step scheduler
- latents = scheduler.step(noise_pred, t, latents).prev_sample
- # need to do CFG?
- # Offload transformer
- transformer.to("cpu")
- torch.cuda.empty_cache()
- # ==========================================
- # Step 3: Decode Latents
- # ==========================================
- print("Decoding latents to video...")
- vae.to(device)
- with torch.no_grad():
- # The Wan VAE requires un-normalizing the latents using its internal config mean/std before decoding
- if hasattr(vae.config, 'latents_mean') and hasattr(vae.config, 'latents_std'):
- mean = torch.tensor(vae.config.latents_mean).view(1, channels, 1, 1, 1).to(device, dtype)
- std = torch.tensor(vae.config.latents_std).view(1, channels, 1, 1, 1).to(device, dtype)
- latents_to_decode = latents * std + mean
- else:
- latents_to_decode = latents / getattr(vae.config, "scaling_factor", 1.0)
- video_tensor = vae.decode(latents_to_decode, return_dict=False)[0]
- # Offload VAE
- vae.to("cpu")
- torch.cuda.empty_cache()
- # ==========================================
- # Step 4: Save Video
- # ==========================================
- print("Saving video...")
- # Squeeze batch dimension[B, C, F, H, W] -> [C, F, H, W]
- video_tensor = video_tensor.squeeze(0).float()
- # Un-normalize from [-1, 1] to [0, 1] then[0, 255]
- video_tensor = (video_tensor / 2 + 0.5).clamp(0, 1)
- video_frames = (video_tensor * 255).to(torch.uint8)
- # torchvision.io.write_video expects dimensions in [Frames, Height, Width, Channels]
- video_frames = video_frames.permute(1, 2, 3, 0).cpu()
- timestamp = int(time.time())
- output_path = f"video_{timestamp}.mp4"
- # Write out the final MP4
- torchvision.io.write_video(output_path, video_frames, fps=24)
- print(f"Done! Saved successfully to '{output_path}'")
Add Comment
Please, Sign In to add comment