Guest User

Untitled

a guest
Apr 17th, 2026
69
0
Never
Not a member of Pastebin yet? Sign Up, it unlocks many cool features!
Python 14.17 KB | None | 0 0
  1. import os
  2. import time
  3. import torch
  4. import torchvision
  5.  
  6. from torch import Tensor
  7. from transformers import AutoTokenizer, T5Gemma2Model, BatchEncoding
  8. from diffusers import AutoencoderKLWan, FlowMatchEulerDiscreteScheduler
  9.  
  10. from transformer.transformer_motif_video import MotifVideoTransformer3DModel
  11.  
  12. # Parameters
  13. 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."
  14. negative_prompt = ""
  15. height = 736
  16. width = 1280
  17. num_frames = 121
  18. num_inference_steps = 50
  19. #cfg = ?
  20. device = "cuda"
  21. dtype = torch.bfloat16
  22. model_path = "./" # Ensure this points to the root of your local Motif-Video-2B repo
  23.  
  24. # Load Components
  25. tokenizer = AutoTokenizer.from_pretrained(model_path, subfolder="tokenizer")
  26. text_encoder = T5Gemma2Model.from_pretrained(model_path, subfolder="text_encoder", torch_dtype=dtype)
  27. vae = AutoencoderKLWan.from_pretrained(model_path, subfolder="vae", torch_dtype=dtype)
  28. transformer = MotifVideoTransformer3DModel.from_pretrained(model_path, subfolder="transformer", torch_dtype=dtype)
  29. scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained(model_path, subfolder="scheduler")
  30.  
  31. text_encoder.eval()
  32. vae.eval()
  33. transformer.eval()
  34.  
  35. # ==========================================
  36. # Util
  37. # ==========================================
  38. # Copied from diffusers.pipelines.flux.pipeline_flux.calculate_shift
  39. '''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):
  40.    m = (max_shift - base_shift) / (max_seq_len - base_seq_len)
  41.    b = base_shift - m * base_seq_len
  42.    mu = image_seq_len * m + b
  43.    return mu
  44.  
  45. # Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion.retrieve_timesteps
  46. 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):
  47.    """
  48.    Retrieve timesteps from the scheduler.
  49.  
  50.    Args:
  51.        scheduler: The noise scheduler to use.
  52.        num_inference_steps: Number of denoising steps.
  53.        device: Device to place timesteps on.
  54.        timesteps: Custom timestep values (mutually exclusive with sigmas).
  55.        sigmas: Custom sigma values (mutually exclusive with timesteps).
  56.        use_linear_quadratic_schedule: If True, use linear-quadratic sigma schedule.
  57.            This overrides the default linear schedule. Requires num_inference_steps
  58.            to be even.
  59.        linear_quadratic_emulating_steps: Controls the linear portion slope.
  60.            Higher values result in gentler slope in the first half. Default: 250.
  61.        **kwargs: Additional arguments passed to scheduler.set_timesteps().
  62.  
  63.    Returns:
  64.        Tuple of (timesteps, num_inference_steps).
  65.  
  66.    Raises:
  67.        ValueError: If both timesteps and sigmas are provided, or if
  68.            use_linear_quadratic_schedule is True but num_inference_steps is odd.
  69.    """
  70.    if timesteps is not None and sigmas is not None:
  71.        raise ValueError("Only one of `timesteps` or `sigmas` can be passed. Please choose one to set custom values")
  72.  
  73.    # Handle linear-quadratic schedule: compute sigmas if flag is set
  74.    if use_linear_quadratic_schedule:
  75.        if sigmas is not None:
  76.            raise ValueError(
  77.                "Cannot use both `sigmas` and `use_linear_quadratic_schedule`. "
  78.                "The linear-quadratic schedule computes sigmas automatically."
  79.            )
  80.        if num_inference_steps is None:
  81.            raise ValueError("`num_inference_steps` must be provided when using `use_linear_quadratic_schedule`.")
  82.        sigmas = get_linear_quadratic_sigmas(
  83.            num_inference_steps=num_inference_steps,
  84.            linear_quadratic_emulating_steps=linear_quadratic_emulating_steps,
  85.        )
  86.  
  87.    if timesteps is not None:
  88.        accepts_timesteps = "timesteps" in set(inspect.signature(scheduler.set_timesteps).parameters.keys())
  89.        if not accepts_timesteps:
  90.            raise ValueError(
  91.                f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom"
  92.                f" timestep schedules. Please check whether you are using the correct scheduler."
  93.            )
  94.        scheduler.set_timesteps(timesteps=timesteps, device=device, **kwargs)
  95.        timesteps = scheduler.timesteps
  96.        num_inference_steps = len(timesteps)
  97.    elif sigmas is not None:
  98.        accept_sigmas = "sigmas" in set(inspect.signature(scheduler.set_timesteps).parameters.keys())
  99.        if not accept_sigmas:
  100.            raise ValueError(
  101.                f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom"
  102.                f" sigmas schedules. Please check whether you are using the correct scheduler."
  103.            )
  104.        scheduler.set_timesteps(sigmas=sigmas, device=device, **kwargs)
  105.        timesteps = scheduler.timesteps
  106.        num_inference_steps = len(timesteps)
  107.    else:
  108.        scheduler.set_timesteps(num_inference_steps, device=device, **kwargs)
  109.        timesteps = scheduler.timesteps
  110.    return timesteps, num_inference_steps
  111.  
  112. def get_linear_quadratic_sigmas(num_inference_steps: int, linear_quadratic_emulating_steps: int = 250) -> np.ndarray:
  113.    """
  114.    Compute a linear-quadratic sigma schedule for flow matching.
  115.  
  116.    This schedule combines:
  117.    - First half: Linear interpolation from high noise to medium noise (slow denoising)
  118.    - Second half: Quadratic interpolation from medium noise to clean (faster denoising)
  119.  
  120.    Convention:
  121.    - sigma=1.0 represents pure noise
  122.    - sigma=0.0 represents clean image
  123.    - Output sigmas are in descending order (1.0 → ~0)
  124.  
  125.    Args:
  126.        num_inference_steps: Total number of denoising steps (must be even).
  127.        linear_quadratic_emulating_steps: Controls the slope of linear interpolation.
  128.            Higher values result in gentler slope in the first half.
  129.  
  130.    Returns:
  131.        np.ndarray: Array of sigma values with shape (num_inference_steps,).
  132.            The scheduler will append a terminal 0.
  133.  
  134.    Raises:
  135.        ValueError: If num_inference_steps is not even.
  136.  
  137.    Reference:
  138.        Linear-quadratic timestep schedule for improved flow matching inference.
  139.    """
  140.    if num_inference_steps % 2 != 0:
  141.        raise ValueError(
  142.            f"num_inference_steps must be even for linear-quadratic schedule, but got {num_inference_steps}"
  143.        )
  144.    
  145.    steps = num_inference_steps
  146.    N = linear_quadratic_emulating_steps
  147.    half_steps = steps // 2
  148.  
  149.    # First half: linear interpolation from 1 toward 0
  150.    # Takes first half_steps values from linspace(1, 0, N+1)
  151.    linear_part = np.linspace(1.0, 0.0, N + 1)[:half_steps]
  152.  
  153.    # Second half: quadratic interpolation
  154.    # Formula: x^2 * (half_steps/N - 1) - (half_steps/N - 1)
  155.    #        = (half_steps/N - 1) * (x^2 - 1)
  156.    # This maps x=0 to (half_steps/N - 1) * (-1) = 1 - half_steps/N
  157.    # and maps x=1 to 0
  158.    x = np.linspace(0.0, 1.0, half_steps + 1)
  159.    scale_factor = half_steps / N - 1  # negative value
  160.    quadratic_part = x**2 * scale_factor - scale_factor
  161.  
  162.    # Concatenate and exclude the last 0 (scheduler appends terminal 0)
  163.    sigmas = np.concatenate([linear_part, quadratic_part])
  164.    sigmas = sigmas[:-1]  # Remove trailing 0, scheduler will append it
  165.  
  166.    return sigmas.astype(np.float32)
  167. '''
  168.  
  169. # ==========================================
  170. # Step 1: Encode Text
  171. # ==========================================
  172. def encode_prompt(prompt, tokenizer, text_encoder, device):
  173.     prompt = [prompt] if isinstance(prompt, str) else prompt
  174.    
  175.     text_inputs = tokenizer(
  176.         prompt,
  177.         padding="max_length",
  178.         max_length=512,
  179.         truncation=True,
  180.         add_special_tokens=True,
  181.         return_attention_mask=True,
  182.         return_tensors="pt",
  183.     )
  184.     #input_ids = text_inputs.input_ids.to(device)
  185.     #attention_mask = text_inputs.attention_mask.to(device)
  186.  
  187.     text_inputs = BatchEncoding(
  188.         {k: v.to(device) if isinstance(v, torch.Tensor) else v for k, v in text_inputs.items()}
  189.     )
  190.     attention_mask = text_inputs.attention_mask.to(device)
  191.  
  192.     '''# Encode text sequence
  193.    # 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
  194.    outputs = text_encoder.encoder(input_ids=input_ids, attention_mask=attention_mask, output_hidden_states=False, return_dict=True)
  195.    
  196.     # T5Gemma2Model returns an encoder_last_hidden_state or last_hidden_state depending on version — check both:
  197.    if hasattr(outputs, "encoder_last_hidden_state") and outputs.encoder_last_hidden_state is not None:
  198.        prompt_embeds = outputs.encoder_last_hidden_state
  199.    else:
  200.        prompt_embeds = outputs.last_hidden_state'''
  201.    
  202.     prompt_embeds = text_encoder.encoder(**text_inputs)[0]
  203.     prompt_embeds = prompt_embeds.to(device=device)
  204.  
  205.     pooled_prompt_embeds = average_pool(prompt_embeds, attention_mask)
  206.    
  207.     return prompt_embeds, attention_mask, pooled_prompt_embeds
  208.  
  209. def average_pool(last_hidden_states: Tensor, attention_mask: Tensor) -> Tensor:
  210.     last_hidden = last_hidden_states.masked_fill(~attention_mask[..., None].bool(), 0.0)
  211.     denom = attention_mask.sum(dim=1, keepdim=True).clamp(min=1)  # avoid div by zero
  212.     return last_hidden.sum(dim=1) / denom
  213.  
  214. print("Encoding prompt...")
  215. text_encoder.to(device)
  216. with torch.no_grad():
  217.     prompt_embeds, attn_mask, pooled_prompt_embeds = encode_prompt(prompt, tokenizer, text_encoder, device)
  218.     neg_prompt_embeds, neg_attn_mask, neg_pooled_prompt_embeds = encode_prompt(negative_prompt, tokenizer, text_encoder, device)
  219.    
  220.     encoder_hidden_states = torch.cat([neg_prompt_embeds, prompt_embeds], dim=0).to(device, dtype)
  221.     encoder_attention_mask = torch.cat([neg_attn_mask, attn_mask], dim=0).to(device, dtype)
  222.     pooled_projections = torch.cat([neg_pooled_prompt_embeds, pooled_prompt_embeds], dim=0).to(device, dtype)
  223.  
  224. text_encoder.to("cpu")
  225. torch.cuda.empty_cache()
  226.  
  227. # ==========================================
  228. # Step 2: Denoising Loop
  229. # ==========================================
  230. transformer.to(device)
  231.  
  232. print("Preparing latents...")
  233.  
  234. # Motif-Video-2B utilizes the Wan 3D VAE structure:
  235. # Temporal downsampling factor is 4, Spatial downsampling factor is 8.
  236. latent_frames = (num_frames - 1) // 4 + 1
  237. latent_height = height // 8
  238. latent_width = width // 8
  239. # VAE z_dim is typically 16 channels
  240. #channels = getattr(transformer.config, "in_channels", 16)
  241. channels = getattr(transformer.config, "in_channels", vae.config.z_dim)  # can read directly?
  242.  
  243. shape = (1, channels, latent_frames, latent_height, latent_width)
  244. latents = torch.randn(shape, generator=torch.Generator(device=device), device=device, dtype=dtype)
  245.  
  246. print("Preparing scheduler...")
  247. scheduler.set_timesteps(num_inference_steps, device=device)
  248. timesteps = scheduler.timesteps
  249.  
  250. print("Starting denoising loop...")
  251. for i, t in enumerate(timesteps):
  252.     print(f"Step {i+1}/{num_inference_steps} (timestep {t.item()})")
  253.    
  254.     with torch.no_grad():
  255.         noise_pred = transformer(
  256.             hidden_states=latents,
  257.             timestep=t.unsqueeze(0).expand(latents.shape[0]), # Expand scalar timestep to 1D tensor
  258.             attention_kwargs=None,
  259.             return_dict=False,
  260.             tread_disabled=True,
  261.             encoder_hidden_states=encoder_hidden_states,
  262.             encoder_attention_mask=encoder_attention_mask,
  263.             pooled_projections=pooled_projections,
  264.             image_embeds=None,
  265.         )[0]
  266.    
  267.     # Step scheduler
  268.     latents = scheduler.step(noise_pred, t, latents).prev_sample
  269.    
  270.     # need to do CFG?
  271.  
  272. # Offload transformer
  273. transformer.to("cpu")
  274. torch.cuda.empty_cache()
  275.  
  276. # ==========================================
  277. # Step 3: Decode Latents
  278. # ==========================================
  279. print("Decoding latents to video...")
  280. vae.to(device)
  281.  
  282. with torch.no_grad():
  283.     # The Wan VAE requires un-normalizing the latents using its internal config mean/std before decoding
  284.     if hasattr(vae.config, 'latents_mean') and hasattr(vae.config, 'latents_std'):
  285.         mean = torch.tensor(vae.config.latents_mean).view(1, channels, 1, 1, 1).to(device, dtype)
  286.         std = torch.tensor(vae.config.latents_std).view(1, channels, 1, 1, 1).to(device, dtype)
  287.         latents_to_decode = latents * std + mean
  288.     else:
  289.         latents_to_decode = latents / getattr(vae.config, "scaling_factor", 1.0)
  290.    
  291.     video_tensor = vae.decode(latents_to_decode, return_dict=False)[0]
  292.  
  293. # Offload VAE
  294. vae.to("cpu")
  295. torch.cuda.empty_cache()
  296.  
  297. # ==========================================
  298. # Step 4: Save Video
  299. # ==========================================
  300. print("Saving video...")
  301. # Squeeze batch dimension[B, C, F, H, W] -> [C, F, H, W]
  302. video_tensor = video_tensor.squeeze(0).float()
  303.  
  304. # Un-normalize from [-1, 1] to [0, 1] then[0, 255]
  305. video_tensor = (video_tensor / 2 + 0.5).clamp(0, 1)
  306. video_frames = (video_tensor * 255).to(torch.uint8)
  307.  
  308. # torchvision.io.write_video expects dimensions in [Frames, Height, Width, Channels]
  309. video_frames = video_frames.permute(1, 2, 3, 0).cpu()
  310.  
  311. timestamp = int(time.time())
  312. output_path = f"video_{timestamp}.mp4"
  313.  
  314. # Write out the final MP4
  315. torchvision.io.write_video(output_path, video_frames, fps=24)
  316. print(f"Done! Saved successfully to '{output_path}'")
Add Comment
Please, Sign In to add comment