mirror of
https://github.com/vladmandic/automatic
synced 2026-09-10 06:48:43 +02:00
Fix loss=nan
This commit is contained in:
@@ -590,7 +590,8 @@ def train_hypernetwork(id_task, hypernetwork_name, learn_rate, batch_size, gradi
|
||||
print(e)
|
||||
|
||||
if shared.cmd_opts.use_ipex:
|
||||
pass
|
||||
shared.sd_model.train()
|
||||
shared.sd_model, optimizer = torch.xpu.optimize(shared.sd_model.to(dtype=torch.float32), optimizer=optimizer, dtype=devices.dtype)
|
||||
else:
|
||||
scaler = torch.cuda.amp.GradScaler()
|
||||
|
||||
|
||||
@@ -430,7 +430,8 @@ 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:
|
||||
pass
|
||||
shared.sd_model.train()
|
||||
shared.sd_model, optimizer = torch.xpu.optimize(shared.sd_model.to(dtype=torch.float32), optimizer=optimizer, dtype=devices.dtype)
|
||||
else:
|
||||
scaler = torch.cuda.amp.GradScaler()
|
||||
|
||||
|
||||
Reference in New Issue
Block a user