fp8_fast mode, cleanup
This commit is contained in:
@@ -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
|
||||
|
||||
|
||||
@@ -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))
|
||||
@@ -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":
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -1,6 +1,4 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import math
|
||||
import torch.distributed as dist
|
||||
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,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,
|
||||
)
|
||||
|
||||
|
||||
@@ -1,7 +1,5 @@
|
||||
import functools
|
||||
import torch.nn as nn
|
||||
from einops import rearrange
|
||||
import torch
|
||||
|
||||
|
||||
def weights_init(m):
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user