Train patches for IPEX

This commit is contained in:
Disty0
2023-06-07 17:25:11 +03:00
parent 2a664b1bdb
commit 3bef3e3eee
3 changed files with 15 additions and 7 deletions
+3
View File
@@ -169,6 +169,9 @@ if args.use_ipex:
CondFunc('torch.nn.modules.GroupNorm.forward',
lambda orig_func, *args, **kwargs: orig_func(args[0], args[1].to(args[0].weight.data.dtype)),
lambda *args, **kwargs: args[2].dtype != args[1].weight.data.dtype)
CondFunc('torch.nn.modules.Linear.forward',
lambda orig_func, *args, **kwargs: orig_func(args[0], args[1].to(args[0].weight.data.dtype)),
lambda *args, **kwargs: args[2].dtype != args[1].weight.data.dtype)
#Use XPU instead of CPU. %20 Perf improvement on weak CPUs.
if args.device_id is not None:
+4 -6
View File
@@ -594,12 +594,10 @@ def train_hypernetwork(id_task, hypernetwork_name, learn_rate, batch_size, gradi
print(e)
if shared.cmd_opts.use_ipex:
scaler = ipex.cpu.autocast._grad_scaler.GradScaler()
#Remove these after Intel adds support for float16 training:
if shared.opts.cuda_dtype == 'BF16':
shared.sd_model = shared.sd_model.bfloat16()
elif shared.opts.cuda_dtype == 'FP32':
shared.sd_model = shared.sd_model.float32()
scaler = ipex.cpu.autocast._grad_scaler.GradScaler() #scaler.step(optimizer): PI_ERROR_INVALID_ARG_VALUE
shared.sd_model = shared.sd_model.to(dtype=torch.float32)
shared.sd_model.train()
shared.sd_model, optimizer = ipex.optimize(shared.sd_model, optimizer=optimizer, dtype=devices.dtype)
else:
scaler = torch.cuda.amp.GradScaler()
@@ -3,6 +3,10 @@ import html
import csv
from collections import namedtuple
import torch
try:
import intel_extension_for_pytorch as ipex # pylint: disable=import-error, unused-import
except:
pass
from tqdm import tqdm
import safetensors.torch
import numpy as np
@@ -433,7 +437,10 @@ def train_embedding(id_task, embedding_name, learn_rate, batch_size, gradient_st
shared.log.info("No saved optimizer exists in checkpoint")
if shared.cmd_opts.use_ipex:
scaler = torch.xpu.amp.GradScaler()
scaler = ipex.cpu.autocast._grad_scaler.GradScaler() #scaler.step(optimizer): PI_ERROR_INVALID_ARG_VALUE
shared.sd_model = shared.sd_model.to(dtype=torch.float32)
shared.sd_model.train()
shared.sd_model, optimizer = ipex.optimize(shared.sd_model, optimizer=optimizer, dtype=devices.dtype)
else:
scaler = torch.cuda.amp.GradScaler()