From dd58129e07ffba858a02348a6b4527cd851a01a8 Mon Sep 17 00:00:00 2001 From: gluttony-10 <52977964+gluttony-10@users.noreply.github.com> Date: Fri, 11 Oct 2024 10:25:45 +0800 Subject: [PATCH] Update pyramid_dit_for_video_gen_pipeline.py Change "torch.float8_e4m3fn, torch.float8_e4m3fn" to "torch.float8_e4m3fn, torch.float8_e5m2". --- pyramid_dit/pyramid_dit_for_video_gen_pipeline.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/pyramid_dit/pyramid_dit_for_video_gen_pipeline.py b/pyramid_dit/pyramid_dit_for_video_gen_pipeline.py index 349975a..3b0ceb8 100644 --- a/pyramid_dit/pyramid_dit_for_video_gen_pipeline.py +++ b/pyramid_dit/pyramid_dit_for_video_gen_pipeline.py @@ -47,7 +47,7 @@ class PyramidDiTForVideoGeneration: ): super().__init__() - if model_dtype in [torch.float8_e4m3fn, torch.float8_e4m3fn]: + if model_dtype in [torch.float8_e4m3fn, torch.float8_e5m2]: self.dtype = torch.bfloat16 else: self.dtype = model_dtype @@ -70,12 +70,12 @@ class PyramidDiTForVideoGeneration: use_temporal_causal=True if not use_flash_attn else False, interp_condition_pos=interp_condition_pos, ) - if model_dtype in [torch.float8_e4m3fn, torch.float8_e4m3fn]: + if model_dtype in [torch.float8_e4m3fn, torch.float8_e5m2]: for name, param in self.dit.named_parameters(): if name != "pos_embedding": param.data = param.data.to(model_dtype) - if model_dtype in [torch.float8_e4m3fn, torch.float8_e4m3fn] and fp8_fastmode: + if model_dtype in [torch.float8_e4m3fn, torch.float8_e5m2] and fp8_fastmode: from ..fp8_optimization import convert_fp8_linear convert_fp8_linear(self.dit, torch.bfloat16)