mirror of
https://github.com/vladmandic/automatic
synced 2026-09-19 17:24:32 +02:00
Train patches for IPEX
This commit is contained in:
@@ -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:
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
Reference in New Issue
Block a user