support fp8 e5m2

This commit is contained in:
kijai
2024-09-07 18:27:02 +03:00
parent cd8e756103
commit 1c70cce963
4 changed files with 24 additions and 8 deletions
+10 -4
View File
@@ -65,10 +65,10 @@ class FluxNetworkTrainer(NetworkTrainer):
)
if args.fp8_base:
# check dtype of model
if model.dtype == torch.float8_e4m3fnuz or model.dtype == torch.float8_e5m2 or model.dtype == torch.float8_e5m2fnuz:
if model.dtype == torch.float8_e4m3fnuz or model.dtype == torch.float8_e5m2fnuz:
raise ValueError(f"Unsupported fp8 model dtype: {model.dtype}")
elif model.dtype == torch.float8_e4m3fn:
logger.info("Loaded fp8 FLUX model")
elif model.dtype == torch.float8_e4m3fn or model.dtype == torch.float8_e5m2:
logger.info(f"Loaded {model.dtype} FLUX model")
if args.split_mode:
model = self.prepare_split_model(model, args, weight_dtype, accelerator)
@@ -115,7 +115,13 @@ class FluxNetworkTrainer(NetworkTrainer):
flux_upper.load_state_dict(sd, strict=False, assign=True)
logger.info("prepare upper model")
target_dtype = torch.float8_e4m3fn if args.fp8_base else weight_dtype
if args.fp8_base:
if args.fp8_dtype and args.fp8_dtype.lower() == "e5m2":
target_dtype = torch.float8_e5m2
else:
target_dtype = torch.float8_e4m3fn
else:
target_dtype =weight_dtype
flux_upper.to(accelerator.device, dtype=target_dtype)
flux_upper.eval()
+9
View File
@@ -3521,6 +3521,13 @@ def add_training_arguments(parser: argparse.ArgumentParser, support_dreambooth:
"--full_bf16", action="store_true", help="bf16 training including gradients / 勾配も含めてbf16で学習する"
) # TODO move to SDXL training, because it is not supported by SD1/2
parser.add_argument("--fp8_base", action="store_true", help="use fp8 for base model / base modelにfp8を使う")
parser.add_argument(
"--fp8_dtype",
type=str,
default="e4m3",
choices=["e4m3", "e5m2"],
help="fp8 dtype selection",
)
parser.add_argument(
"--ddp_timeout",
@@ -4782,6 +4789,8 @@ def prepare_dtype(args: argparse.Namespace):
save_dtype = torch.float32
elif args.save_precision == "fp8_e4m3fn":
save_dtype = torch.float8_e4m3fn
elif args.save_precision == "fp8_e5m2":
save_dtype = torch.float8_e5m2
return weight_dtype, save_dtype
+2 -2
View File
@@ -303,7 +303,7 @@ class InitFluxLoRATraining:
"highvram": ("BOOLEAN", {"default": False, "tooltip": "memory mode"}),
"fp8_base": ("BOOLEAN", {"default": True, "tooltip": "use fp8 for base model"}),
"gradient_dtype": (["fp32", "fp16", "bf16"], {"default": "fp32", "tooltip": "the actual dtype training uses"}),
"save_dtype": (["fp32", "fp16", "bf16", "fp8_e4m3fn"], {"default": "bf16", "tooltip": "the dtype to save checkpoints as"}),
"save_dtype": (["fp32", "fp16", "bf16", "fp8_e4m3fn", "fp8_e5m2"], {"default": "bf16", "tooltip": "the dtype to save checkpoints as"}),
"attention_mode": (["sdpa", "xformers", "disabled"], {"default": "sdpa", "tooltip": "memory efficient attention mode"}),
"sample_prompts": ("STRING", {"multiline": True, "default": "illustration of a kitten | photograph of a turtle", "tooltip": "validation sample prompts, for multiple prompts, separate by `|`"}),
},
@@ -1414,7 +1414,7 @@ class ExtractFluxLoRA:
"finetuned_model": (folder_paths.get_filename_list("unet"), ),
"output_path": ("STRING", {"default": f"{str(os.path.join(folder_paths.models_dir, 'loras', 'Flux'))}"}),
"dim": ("INT", {"default": 4, "min": 2, "max": 1024, "step": 2, "tooltip": "LoRA rank"}),
"save_dtype": (["fp32", "fp16", "bf16", "fp8_e4m3fn"], {"default": "bf16", "tooltip": "the dtype to save the LoRA as"}),
"save_dtype": (["fp32", "fp16", "bf16", "fp8_e4m3fn", "fp8_e5m2"], {"default": "bf16", "tooltip": "the dtype to save the LoRA as"}),
"load_device": (["cpu", "cuda"], {"default": "cuda", "tooltip": "the device to load the model to"}),
"store_device": (["cpu", "cuda"], {"default": "cpu", "tooltip": "the device to store the LoRA as"}),
"clamp_quantile": ("FLOAT", {"default": 0.99, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "clamp quantile"}),
+3 -2
View File
@@ -555,11 +555,12 @@ class NetworkTrainer:
args.mixed_precision != "no"
), "fp8_base requires mixed precision='fp16' or 'bf16'"
accelerator.print("enable fp8 training for U-Net.")
unet_weight_dtype = torch.float8_e4m3fn
unet_weight_dtype = torch.float8_e4m3fn if args.fp8_dtype == "e4m3" else torch.float8_e5m2
accelerator.print(f"unet_weight_dtype: {unet_weight_dtype}")
if not args.fp8_base_unet and not args.network_train_unet_only:
accelerator.print("enable fp8 training for Text Encoder.")
te_weight_dtype = weight_dtype if args.fp8_base_unet else torch.float8_e4m3fn
te_weight_dtype = torch.float8_e4m3fn if args.fp8_dtype == "e4m3" else torch.float8_e5m2
# unet.to(accelerator.device) # this makes faster `to(dtype)` below, but consumes 23 GB VRAM
# unet.to(dtype=unet_weight_dtype) # without moving to gpu, this takes a lot of time and main memory