fp8_fast mode, cleanup

This commit is contained in:
Jukka Seppänen
2024-10-11 00:47:08 +03:00
parent 51fb1da7ee
commit 5252379284
11 changed files with 60 additions and 61 deletions
@@ -5,8 +5,7 @@ import numpy as np
import torch
from diffusers.configuration_utils import ConfigMixin, register_to_config
from diffusers.utils import BaseOutput, logging
from diffusers.utils.torch_utils import randn_tensor
from diffusers.utils import BaseOutput
from diffusers.schedulers.scheduling_utils import SchedulerMixin
#from IPython import embed
+45
View File
@@ -0,0 +1,45 @@
#based on ComfyUI's and MinusZoneAI's fp8_linear optimization
import torch
import torch.nn as nn
def fp8_linear_forward(cls, original_dtype, input):
weight_dtype = cls.weight.dtype
if weight_dtype in [torch.float8_e4m3fn, torch.float8_e5m2]:
if len(input.shape) == 3:
if weight_dtype == torch.float8_e4m3fn:
inn = input.reshape(-1, input.shape[2]).to(torch.float8_e5m2)
else:
inn = input.reshape(-1, input.shape[2]).to(torch.float8_e4m3fn)
w = cls.weight.t()
scale_weight = torch.ones((1), device=input.device, dtype=torch.float32)
scale_input = scale_weight
bias = cls.bias.to(original_dtype) if cls.bias is not None else None
out_dtype = original_dtype
if bias is not None:
o = torch._scaled_mm(inn, w, out_dtype=out_dtype, bias=bias, scale_a=scale_input, scale_b=scale_weight)
else:
o = torch._scaled_mm(inn, w, out_dtype=out_dtype, scale_a=scale_input, scale_b=scale_weight)
if isinstance(o, tuple):
o = o[0]
return o.reshape((-1, input.shape[1], cls.weight.shape[0]))
else:
cls.to(original_dtype)
out = cls.original_forward(input.to(original_dtype))
cls.to(original_dtype)
return out
else:
return cls.original_forward(input)
def convert_fp8_linear(module, original_dtype):
setattr(module, "fp8_matmul_enabled", True)
for name, module in module.named_modules():
if isinstance(module, nn.Linear):
original_forward = module.forward
setattr(module, "original_forward", original_forward)
setattr(module, "forward", lambda input, m=module: fp8_linear_forward(m, original_dtype, input))
+4 -22
View File
@@ -37,7 +37,7 @@ class DownloadAndLoadPyramidFlowModel:
"text_encoder_dtype": (["fp16", "fp32", "bf16"],{"default": "bf16", }),
"vae_dtype": (["fp16", "fp32", "bf16"],{"default": "bf16", }),
"use_flash_attn": ("BOOLEAN", {"default": False}),
#"fp8_transformer": (['disabled', 'enabled', 'fastmode'], {"default": 'disabled', "tooltip": "enabled casts the transformer to torch.float8_e4m3fn, fastmode is only for latest nvidia GPUs"}),
"fp8_fastmode": ("BOOLEAN",{"default": False, "tooltip": "fastmode is only for latest nvidia GPUs"}),
#"compile": (["disabled","onediff","torch"], {"tooltip": "compile the model for faster inference, these are advanced options only available on Linux, see readme for more info"}),
}
}
@@ -47,7 +47,7 @@ class DownloadAndLoadPyramidFlowModel:
FUNCTION = "loadmodel"
CATEGORY = "PyramidFlowWrapper"
def loadmodel(self, model, variant, model_dtype, text_encoder_dtype, vae_dtype, use_flash_attn=False):
def loadmodel(self, model, variant, model_dtype, text_encoder_dtype, vae_dtype, fp8_fastmode, use_flash_attn=False):
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
@@ -79,26 +79,8 @@ class DownloadAndLoadPyramidFlowModel:
vae_dtype,
model_variant=variant,
use_flash_attn=use_flash_attn,
)
# #fp8
# if fp8_transformer == "enabled" or fp8_transformer == "fastmode":
# if "2b" in model:
# for name, param in transformer.named_parameters():
# if name != "pos_embedding":
# param.data = param.data.to(torch.float8_e4m3fn)
# elif "I2V" in model:
# for name, param in transformer.named_parameters():
# if "patch_embed" not in name:
# param.data = param.data.to(torch.float8_e4m3fn)
# else:
# transformer.to(torch.float8_e4m3fn)
# if fp8_transformer == "fastmode":
# from .fp8_optimization import convert_fp8_linear
# convert_fp8_linear(transformer, dtype)
fp8_fastmode=fp8_fastmode,
)
# # compilation
# if compile == "torch":
-2
View File
@@ -1,6 +1,5 @@
import torch
import torch.nn as nn
import os
import torch.nn.functional as F
from einops import rearrange
@@ -9,7 +8,6 @@ from diffusers.models.modeling_utils import ModelMixin
from diffusers.configuration_utils import ConfigMixin, register_to_config
from diffusers.utils import is_torch_version
from typing import Any, Callable, Dict, List, Optional, Union
from tqdm import tqdm
from .modeling_embedding import PatchEmbed3D, CombinedTimestepConditionEmbeddings
from .modeling_normalization import AdaLayerNormContinuous
@@ -8,12 +8,9 @@ from einops import rearrange
from diffusers.utils.torch_utils import randn_tensor
import math
import PIL
from PIL import Image
from tqdm import tqdm
from torchvision import transforms
from copy import deepcopy
from typing import Any, Callable, Dict, List, Optional, Union
from typing import List, Optional, Union
from ..diffusion_schedulers import PyramidFlowMatchEulerDiscreteScheduler
from ..video_vae.modeling_causal_vae import CausalVideoVAE
@@ -46,7 +43,7 @@ class PyramidDiTForVideoGeneration:
model_variant="diffusion_transformer_768p", timestep_shift=1.0, stage_range=[0, 1/3, 2/3, 1],
sample_ratios=[1, 1, 1], scheduler_gamma=1/3, use_flash_attn=False,
load_text_encoder=True, load_vae=True, max_temporal_length=31, frame_per_unit=1, use_temporal_causal=True,
corrupt_ratio=1/3, interp_condition_pos=True, stages=[1, 2, 4], **kwargs,
corrupt_ratio=1/3, interp_condition_pos=True, stages=[1, 2, 4], fp8_fastmode=False, **kwargs,
):
super().__init__()
@@ -78,6 +75,10 @@ class PyramidDiTForVideoGeneration:
if name != "pos_embedding":
param.data = param.data.to(model_dtype)
if fp8_fastmode == "fastmode":
from ..fp8_optimization import convert_fp8_linear
convert_fp8_linear(self.dit, torch.bfloat16)
# The text encoder
if load_text_encoder:
self.text_encoder = SD3TextEncoderWithMask(model_path, torch_dtype=text_encoder_dtype)
-2
View File
@@ -1,6 +1,4 @@
import torch
import torch.nn as nn
import math
import torch.distributed as dist
+2 -7
View File
@@ -1,17 +1,12 @@
import io
import os
import math
import time
import json
import glob
from collections import defaultdict, deque, OrderedDict
from collections import defaultdict, deque
import datetime
import numpy as np
from pathlib import Path
import argparse
import torch
from torch import optim as optim
import torch.distributed as dist
-2
View File
@@ -13,9 +13,7 @@
# limitations under the License.
from typing import Any, Dict, Optional, Tuple, Union
import numpy as np
import torch
import torch.nn.functional as F
from torch import nn
from einops import rearrange
+1 -8
View File
@@ -1,25 +1,18 @@
from typing import Tuple, Union
import torch
import torch.nn as nn
from torch.utils.checkpoint import checkpoint
import torch.nn.functional as F
from collections import deque
from einops import rearrange
from timm.models.layers import trunc_normal_
#from IPython import embed
from torch import Tensor
from ..utils import (
is_context_parallel_initialized,
get_context_parallel_group,
get_context_parallel_world_size,
get_context_parallel_rank,
get_context_parallel_group_rank,
)
)
from .context_parallel_ops import (
conv_scatter_to_context_parallel_region,
conv_gather_from_context_parallel_region,
cp_pass_from_previous_rank,
)
-2
View File
@@ -1,7 +1,5 @@
import functools
import torch.nn as nn
from einops import rearrange
import torch
def weights_init(m):
-8
View File
@@ -17,11 +17,9 @@ from typing import Optional, Tuple
import numpy as np
import torch
import torch.nn as nn
from einops import rearrange
from diffusers.utils import BaseOutput, is_torch_version
from diffusers.utils.torch_utils import randn_tensor
from diffusers.models.attention_processor import SpatialNorm
from .modeling_block import (
UNetMidBlock2D,
CausalUNetMidBlock2D,
@@ -30,12 +28,6 @@ from .modeling_block import (
get_input_layer,
get_output_layer,
)
from .modeling_resnet import (
Downsample2D,
Upsample2D,
TemporalDownsample2x,
TemporalUpsample2x,
)
from .modeling_causal_conv import CausalConv3d, CausalGroupNorm