fix parallel inference for Helios-Mid & Helios-Distilled

This commit is contained in:
SHYuanBest
2026-03-08 06:01:44 +00:00
parent cb524be214
commit 7903ec31b5
3 changed files with 6 additions and 0 deletions
@@ -21,6 +21,7 @@ import numpy as np
import regex as re
import torch
import torch.nn.functional as F
from accelerate.utils import broadcast
from transformers import AutoTokenizer, UMT5EncoderModel
from diffusers.callbacks import MultiPipelineCallbacks, PipelineCallback
@@ -697,6 +698,7 @@ class HeliosPipeline(DiffusionPipeline, HeliosLoraLoaderMixin):
batch_size, channel, num_frames, height, width = latents.shape
noise = self.sample_block_noise(batch_size, channel, num_frames, height, width, patch_size, device)
noise = noise.to(device=device, dtype=transformer_dtype)
noise = broadcast(noise, from_process=0)
latents = alpha * latents + beta * noise # To fix the block artifact
if self.config.is_distilled:
+2
View File
@@ -21,6 +21,7 @@ from typing import Any, Callable, Dict, List, Optional, Union
import regex as re
import torch
import torch.nn.functional as F
from accelerate.utils import broadcast
from einops import rearrange
from transformers import AutoTokenizer, UMT5EncoderModel
@@ -670,6 +671,7 @@ class HeliosPipeline(DiffusionPipeline, WanLoraLoaderMixin):
batch_size, channel, num_frames, height, width = latents.shape
noise = self.sample_block_noise(batch_size, channel, num_frames, height, width)
noise = noise.to(device=device, dtype=transformer_dtype)
noise = broadcast(noise, from_process=0)
latents = alpha * latents + beta * noise # To fix the block artifact
if use_dmd:
+2
View File
@@ -21,6 +21,7 @@ from typing import Any, Callable, Dict, List, Optional, Union
import regex as re
import torch
import torch.nn.functional as F
from accelerate.utils import broadcast
from einops import rearrange
from transformers import AutoTokenizer, UMT5EncoderModel
@@ -677,6 +678,7 @@ class HeliosPipeline(DiffusionPipeline, WanLoraLoaderMixin):
batch_size, channel, num_frames, height, width = latents.shape
noise = self.sample_block_noise(batch_size, channel, num_frames, height, width)
noise = noise.to(device=device, dtype=transformer_dtype)
noise = broadcast(noise, from_process=0)
latents = alpha * latents + beta * noise # To fix the block artifact
if use_dmd: