[diffusion] improve: avoid cosmos3 cpu float video postprocess (#26956)
This commit is contained in:
+14
-13
@@ -822,11 +822,6 @@ class Cosmos3DecodingStage(PipelineStage):
|
|||||||
self._latents_mean = None
|
self._latents_mean = None
|
||||||
self._latents_std = None
|
self._latents_std = None
|
||||||
self._guardrails = guardrails
|
self._guardrails = guardrails
|
||||||
# Use VideoProcessor for postprocessing (same as other video pipelines)
|
|
||||||
from diffusers.video_processor import VideoProcessor
|
|
||||||
|
|
||||||
vae_scale_factor = getattr(vae.config, "scale_factor_spatial", 16)
|
|
||||||
self.video_processor = VideoProcessor(vae_scale_factor=vae_scale_factor)
|
|
||||||
if guardrails:
|
if guardrails:
|
||||||
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.cosmos3_guardrails import (
|
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.cosmos3_guardrails import (
|
||||||
_init_guardrails,
|
_init_guardrails,
|
||||||
@@ -876,6 +871,16 @@ class Cosmos3DecodingStage(PipelineStage):
|
|||||||
|
|
||||||
return video
|
return video
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _postprocess_tensor(decoded: torch.Tensor) -> torch.Tensor:
|
||||||
|
return (decoded * 0.5 + 0.5).clamp(0, 1).float()
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _postprocess_video_np(video: torch.Tensor, is_image_gen: bool) -> np.ndarray:
|
||||||
|
if is_image_gen:
|
||||||
|
return video.squeeze(2).permute(0, 2, 3, 1).cpu().numpy()
|
||||||
|
return video.permute(0, 2, 3, 4, 1).cpu().numpy()
|
||||||
|
|
||||||
def forward(self, batch: Req, server_args: ServerArgs):
|
def forward(self, batch: Req, server_args: ServerArgs):
|
||||||
"""Decode latents to video, or to a single image for T2I."""
|
"""Decode latents to video, or to a single image for T2I."""
|
||||||
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import (
|
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import (
|
||||||
@@ -900,25 +905,21 @@ class Cosmos3DecodingStage(PipelineStage):
|
|||||||
self.vae.to("cpu", non_blocking=True)
|
self.vae.to("cpu", non_blocking=True)
|
||||||
|
|
||||||
self.log_info(f"Decoded tensor shape: {decoded.shape}")
|
self.log_info(f"Decoded tensor shape: {decoded.shape}")
|
||||||
|
output = self._postprocess_tensor(decoded)
|
||||||
if is_image_gen:
|
|
||||||
output = self.video_processor.postprocess(
|
|
||||||
decoded.squeeze(2), output_type="np"
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
output = self.video_processor.postprocess_video(decoded, output_type="np")
|
|
||||||
self.log_info(f"Postprocessed video shape: {output.shape}")
|
|
||||||
|
|
||||||
if self._guardrails and batch.use_guardrails is not False:
|
if self._guardrails and batch.use_guardrails is not False:
|
||||||
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.cosmos3_guardrails import (
|
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.cosmos3_guardrails import (
|
||||||
check_video_safety,
|
check_video_safety,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
output = self._postprocess_video_np(output, is_image_gen)
|
||||||
if is_image_gen:
|
if is_image_gen:
|
||||||
# check_video_safety expects [B, T, H, W, C]; wrap then unwrap.
|
# check_video_safety expects [B, T, H, W, C]; wrap then unwrap.
|
||||||
output = check_video_safety(output[:, np.newaxis, ...])[:, 0, ...]
|
output = check_video_safety(output[:, np.newaxis, ...])[:, 0, ...]
|
||||||
else:
|
else:
|
||||||
output = check_video_safety(output)
|
output = check_video_safety(output)
|
||||||
|
elif not is_image_gen:
|
||||||
|
self.log_info(f"Postprocessed video tensor shape: {output.shape}")
|
||||||
|
|
||||||
return OutputBatch(
|
return OutputBatch(
|
||||||
output=output,
|
output=output,
|
||||||
|
|||||||
Reference in New Issue
Block a user