Version 3 / rewrite
This commit is contained in:
+4
-1
@@ -1,12 +1,15 @@
|
|||||||
raw/
|
raw/
|
||||||
images/
|
images/
|
||||||
latent_*/
|
latents/
|
||||||
|
latents
|
||||||
vae/
|
vae/
|
||||||
models/
|
models/
|
||||||
|
other/
|
||||||
test.py
|
test.py
|
||||||
*.png
|
*.png
|
||||||
*.zip
|
*.zip
|
||||||
*.npy
|
*.npy
|
||||||
|
*.pth
|
||||||
*.ckpt
|
*.ckpt
|
||||||
*.safetensors
|
*.safetensors
|
||||||
|
|
||||||
|
|||||||
+47
-22
@@ -2,30 +2,55 @@ import torch
|
|||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
import numpy as np
|
import numpy as np
|
||||||
|
|
||||||
|
class Block(nn.Module):
|
||||||
|
def __init__(self, size):
|
||||||
|
super().__init__()
|
||||||
|
self.join = nn.ReLU()
|
||||||
|
self.long = nn.Sequential(
|
||||||
|
nn.Conv2d(size, size, kernel_size=3, stride=1, padding=1),
|
||||||
|
nn.LeakyReLU(0.1),
|
||||||
|
nn.Conv2d(size, size, kernel_size=3, stride=1, padding=1),
|
||||||
|
nn.LeakyReLU(0.1),
|
||||||
|
nn.Conv2d(size, size, kernel_size=3, stride=1, padding=1),
|
||||||
|
nn.Dropout(0.2)
|
||||||
|
)
|
||||||
|
def forward(self, x):
|
||||||
|
y = self.long(x)
|
||||||
|
z = self.join(y + x)
|
||||||
|
return z
|
||||||
|
|
||||||
class Interposer(nn.Module):
|
class Interposer(nn.Module):
|
||||||
def __init__(self):
|
def __init__(self):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
|
self.chan = 4 # in/out channels
|
||||||
# it looks like a spaceship if you squint :D
|
self.hid = 128
|
||||||
module_list = [
|
|
||||||
#############)
|
|
||||||
#############)
|
|
||||||
#||#
|
|
||||||
#||#
|
|
||||||
nn.Conv2d(4, 32, kernel_size=5, padding=2),
|
|
||||||
nn.ReLU(),
|
|
||||||
nn.Conv2d(32, 128, kernel_size=7, padding=3),
|
|
||||||
nn.ReLU(),
|
|
||||||
nn.Conv2d(128, 32, kernel_size=7, padding=3),
|
|
||||||
nn.ReLU(),
|
|
||||||
nn.Conv2d(32, 4, kernel_size=5, padding=2),
|
|
||||||
#||#
|
|
||||||
#||#
|
|
||||||
#############)
|
|
||||||
#############)
|
|
||||||
]
|
|
||||||
|
|
||||||
self.sequential = nn.Sequential(*module_list)
|
# expand channels
|
||||||
|
self.head_join = nn.ReLU()
|
||||||
|
self.head_short = nn.Conv2d(self.chan, self.hid, kernel_size=3, stride=1, padding=1)
|
||||||
|
self.head_long = nn.Sequential(
|
||||||
|
nn.Conv2d(self.chan, self.hid, kernel_size=3, stride=1, padding=1),
|
||||||
|
nn.LeakyReLU(0.1),
|
||||||
|
nn.Conv2d(self.hid, self.hid, kernel_size=3, stride=1, padding=1),
|
||||||
|
nn.LeakyReLU(0.1),
|
||||||
|
nn.Conv2d(self.hid, self.hid, kernel_size=3, stride=1, padding=1),
|
||||||
|
)
|
||||||
|
# not sure if this is how residuals work
|
||||||
|
self.core = nn.Sequential(
|
||||||
|
Block(self.hid),
|
||||||
|
Block(self.hid),
|
||||||
|
Block(self.hid),
|
||||||
|
)
|
||||||
|
# reduce channels
|
||||||
|
self.tail = nn.Sequential(
|
||||||
|
nn.ReLU(),
|
||||||
|
nn.Conv2d(self.hid, self.chan, kernel_size=3, stride=1, padding=1)
|
||||||
|
)
|
||||||
|
|
||||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
def forward(self, x):
|
||||||
return self.sequential(x)
|
y = self.head_join(
|
||||||
|
self.head_long(x)+
|
||||||
|
self.head_short(x)
|
||||||
|
)
|
||||||
|
z = self.core(y)
|
||||||
|
return self.tail(z)
|
||||||
|
|||||||
+1
-1
@@ -15,7 +15,7 @@ def process_lines(lines):
|
|||||||
[int(x[0]) for x in vals],
|
[int(x[0]) for x in vals],
|
||||||
[math.log(float(x[1])) for x in vals],
|
[math.log(float(x[1])) for x in vals],
|
||||||
)
|
)
|
||||||
if len(vals[0]) == 3:
|
if len(vals[0]) >= 3:
|
||||||
eval_loss[name] = (
|
eval_loss[name] = (
|
||||||
[int(x[0]) for x in vals],
|
[int(x[0]) for x in vals],
|
||||||
[math.log(float(x[2])) for x in vals],
|
[math.log(float(x[2])) for x in vals],
|
||||||
|
|||||||
@@ -1,65 +0,0 @@
|
|||||||
import os
|
|
||||||
import hashlib
|
|
||||||
import argparse
|
|
||||||
from tqdm import tqdm
|
|
||||||
from PIL import Image
|
|
||||||
from queue import Queue
|
|
||||||
from threading import Thread
|
|
||||||
|
|
||||||
if not os.path.isdir("images"):
|
|
||||||
os.mkdir("images")
|
|
||||||
|
|
||||||
def parse_args():
|
|
||||||
parser = argparse.ArgumentParser(description="Preprocess images")
|
|
||||||
parser.add_argument("-r", "--res", type=int, default=768, help="Target resolution")
|
|
||||||
parser.add_argument("-t", "--threads", type=int, default=4, help="No. of CPU threads to use")
|
|
||||||
parser.add_argument('--src', default="raw", help="Source folder with images")
|
|
||||||
return parser.parse_args()
|
|
||||||
|
|
||||||
def process(fname, folder, resolution):
|
|
||||||
src = os.path.join(folder, fname)
|
|
||||||
md5 = hashlib.md5(open(src,'rb').read()).hexdigest()
|
|
||||||
out = os.path.join("images", f"{md5}.png")
|
|
||||||
if os.path.isfile(out):
|
|
||||||
return
|
|
||||||
img = Image.open(src)
|
|
||||||
img = img.convert('RGB')
|
|
||||||
target = (resolution, resolution)
|
|
||||||
if min(img.height, img.width) < 256:
|
|
||||||
return
|
|
||||||
|
|
||||||
if img.width > img.height:
|
|
||||||
target = (int(img.width/img.height*resolution), resolution)
|
|
||||||
elif img.height > img.width:
|
|
||||||
target = (resolution, int(img.height/img.width*resolution))
|
|
||||||
img = img.resize(target, Image.LANCZOS)
|
|
||||||
img = img.crop([0,0,resolution,resolution])
|
|
||||||
img.save(out)
|
|
||||||
|
|
||||||
def thread(queue, pbar, folder, resolution):
|
|
||||||
while not queue.empty():
|
|
||||||
fname = queue.get()
|
|
||||||
try: process(fname, folder, resolution)
|
|
||||||
except: pass
|
|
||||||
queue.task_done()
|
|
||||||
pbar.update()
|
|
||||||
|
|
||||||
args = parse_args()
|
|
||||||
files = os.listdir(args.src)
|
|
||||||
pbar = tqdm(total=len(files),unit="img")
|
|
||||||
queue = Queue()
|
|
||||||
[queue.put(x) for x in files]
|
|
||||||
|
|
||||||
for _ in range(args.threads):
|
|
||||||
Thread(
|
|
||||||
target=thread,
|
|
||||||
args=(
|
|
||||||
queue,
|
|
||||||
pbar,
|
|
||||||
args.src,
|
|
||||||
args.res,
|
|
||||||
),
|
|
||||||
daemon=True,
|
|
||||||
).start()
|
|
||||||
|
|
||||||
queue.join()
|
|
||||||
@@ -1,51 +0,0 @@
|
|||||||
import os
|
|
||||||
import torch
|
|
||||||
import numpy as np
|
|
||||||
from torchvision import transforms
|
|
||||||
from diffusers import AutoencoderKL
|
|
||||||
from tqdm import tqdm
|
|
||||||
from PIL import Image
|
|
||||||
|
|
||||||
from vae import get_vae
|
|
||||||
|
|
||||||
def encode(vae, img):
|
|
||||||
"""image [PIL Image] -> latent [np array]"""
|
|
||||||
inp = transforms.ToTensor()(img).unsqueeze(0)
|
|
||||||
inp = inp.to("cuda") # move to GPU
|
|
||||||
latent = vae.encode(inp*2.0-1.0)
|
|
||||||
latent = latent.latent_dist.sample()
|
|
||||||
return latent.cpu().detach()
|
|
||||||
|
|
||||||
def process_folder(vae, v):
|
|
||||||
if not os.path.isdir(f"latent_{v}"):
|
|
||||||
os.mkdir(f"latent_{v}")
|
|
||||||
|
|
||||||
vae.to("cuda")
|
|
||||||
for i in tqdm(os.listdir("images")):
|
|
||||||
src = os.path.join("images", i)
|
|
||||||
img = Image.open(src)
|
|
||||||
dst = os.path.join(f"latent_{v}", f"{os.path.splitext(i)[0]}.npy")
|
|
||||||
latent = encode(vae, img)
|
|
||||||
np.save(dst, latent)
|
|
||||||
vae.to("cpu")
|
|
||||||
|
|
||||||
def run_v1(file_path=None):
|
|
||||||
vae = get_vae("v1", file_path)
|
|
||||||
process_folder(vae, "v1")
|
|
||||||
del vae
|
|
||||||
|
|
||||||
def run_v2(file_path=None):
|
|
||||||
vae = get_vae("v2", file_path)
|
|
||||||
process_folder(vae, "v2")
|
|
||||||
del vae
|
|
||||||
|
|
||||||
def run_xl(file_path=None):
|
|
||||||
vae = get_vae("xl", file_path)
|
|
||||||
process_folder(vae, "xl")
|
|
||||||
del vae
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
|
||||||
# run_v1("./vae/ft-mse-840000.ckpt") # probably doesn't reflect internal SD latent
|
|
||||||
run_v1()
|
|
||||||
# run_v2() # v2 and v1 share a latent space
|
|
||||||
run_xl("./vae/sdxl_v0.9.safetensors") # 1.0 has artifacts
|
|
||||||
@@ -3,23 +3,27 @@ import torch
|
|||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import argparse
|
import argparse
|
||||||
import random
|
|
||||||
from PIL import Image
|
from PIL import Image
|
||||||
from tqdm import tqdm
|
from tqdm import tqdm
|
||||||
from safetensors.torch import save_file, load_file
|
from safetensors.torch import save_file, load_file
|
||||||
|
from torch.utils.data import DataLoader, Dataset
|
||||||
|
|
||||||
from interposer import Interposer
|
from interposer import Interposer
|
||||||
from vae import get_vae
|
from vae import get_vae
|
||||||
|
|
||||||
|
torch.backends.cudnn.benchmark = True
|
||||||
|
torch.manual_seed(0)
|
||||||
|
|
||||||
def parse_args():
|
def parse_args():
|
||||||
parser = argparse.ArgumentParser(description="Train latent interposer model")
|
parser = argparse.ArgumentParser(description="Train latent interposer model")
|
||||||
parser.add_argument("--steps", type=int, default=500000, help="No. of training steps")
|
parser.add_argument("--steps", type=int, default=500000, help="No. of training steps")
|
||||||
parser.add_argument('--bs', type=int, default=1, help="Batch size")
|
parser.add_argument('--bs', type=int, default=4, help="Batch size")
|
||||||
parser.add_argument('--lr', default="1e-8", help="Learning rate")
|
parser.add_argument('--lr', default="1e-4", help="Learning rate")
|
||||||
parser.add_argument("-n", "--save_every_n", type=int, dest="save", default=50000, help="Save model/sample periodically")
|
parser.add_argument("-n", "--save_every_n", type=int, dest="save", default=50000, help="Save model/sample periodically")
|
||||||
parser.add_argument('--src', choices=["v1","xl"], required=True, help="Source latent format")
|
parser.add_argument('--src', choices=["v1","xl"], required=True, help="Source latent format")
|
||||||
parser.add_argument('--dst', choices=["v1","xl"], required=True, help="Destination latent format")
|
parser.add_argument('--dst', choices=["v1","xl"], required=True, help="Destination latent format")
|
||||||
parser.add_argument('--resume', help="Checkpoint to resume from")
|
parser.add_argument('--resume', help="Checkpoint to resume from")
|
||||||
|
parser.add_argument('--cosine', action=argparse.BooleanOptionalAction, help="Use cosine scheduler to taper off LR")
|
||||||
args = parser.parse_args()
|
args = parser.parse_args()
|
||||||
if args.src == args.dst:
|
if args.src == args.dst:
|
||||||
parser.error("--src and --dst can't be the same")
|
parser.error("--src and --dst can't be the same")
|
||||||
@@ -29,25 +33,6 @@ def parse_args():
|
|||||||
parser.error("--lr must be a valid float eg. 0.001 or 1e-3")
|
parser.error("--lr must be a valid float eg. 0.001 or 1e-3")
|
||||||
return args
|
return args
|
||||||
|
|
||||||
class Latent:
|
|
||||||
def __init__(self, md5, lat_src, lat_dst, dev):
|
|
||||||
if lat_src == "v1": src = os.path.join("latent_v1", f"{md5}.npy")
|
|
||||||
if lat_src == "xl": src = os.path.join("latent_xl", f"{md5}.npy")
|
|
||||||
|
|
||||||
if lat_dst == "v1": dst = os.path.join("latent_v1", f"{md5}.npy")
|
|
||||||
if lat_dst == "xl": dst = os.path.join("latent_xl", f"{md5}.npy")
|
|
||||||
|
|
||||||
self.src = torch.from_numpy(np.load(src)).to(dev)
|
|
||||||
self.dst = torch.from_numpy(np.load(dst)).to(dev)
|
|
||||||
|
|
||||||
def load_latents(src, dst, dev):
|
|
||||||
print("Loading latents from disk")
|
|
||||||
latents = []
|
|
||||||
for i in tqdm(os.listdir("images")):
|
|
||||||
md5 = os.path.splitext(i)[0]
|
|
||||||
latents.append(Latent(md5, src, dst, dev))
|
|
||||||
return latents
|
|
||||||
|
|
||||||
vae = None
|
vae = None
|
||||||
def sample_decode(latent, filename, version):
|
def sample_decode(latent, filename, version):
|
||||||
global vae
|
global vae
|
||||||
@@ -66,66 +51,147 @@ def sample_decode(latent, filename, version):
|
|||||||
out = Image.fromarray(out)
|
out = Image.fromarray(out)
|
||||||
out.save(filename)
|
out.save(filename)
|
||||||
|
|
||||||
|
def get_eval_data(dataset, src_path, dst_path, target_dev):
|
||||||
|
if os.path.isfile(src_path) and os.path.isfile(dst_path):
|
||||||
|
src = LatentDataset.load_latent(None, src_path)
|
||||||
|
dst = LatentDataset.load_latent(None, dst_path)
|
||||||
|
else:
|
||||||
|
src = dataset[0][0]
|
||||||
|
dst = dataset[0][1]
|
||||||
|
src = src.float().to(target_dev).unsqueeze(0)
|
||||||
|
dst = dst.float().to(target_dev).unsqueeze(0)
|
||||||
|
return(src, dst)
|
||||||
|
|
||||||
|
def eval_model(step, model, criterion, scheduler, src, dst):
|
||||||
|
with torch.no_grad():
|
||||||
|
t_pred = model(src)
|
||||||
|
t_loss = criterion(t_pred, dst)
|
||||||
|
tqdm.write(f"{str(step):<10} {loss.data.item():.4e}|{t_loss.data.item():.4e} @ {float(scheduler.get_last_lr()[0]):.4e}")
|
||||||
|
log.write(f"{step},{loss.data.item()},{t_loss.data.item()},{float(scheduler.get_last_lr()[0])}\n")
|
||||||
|
log.flush()
|
||||||
|
|
||||||
|
def save_model(step, model, optim, lat, src, dst):
|
||||||
|
with torch.no_grad():
|
||||||
|
out = model(lat)
|
||||||
|
output_name = f"./models/{src}-to-{dst}_interposer_e{round(step/1000)}k"
|
||||||
|
sample_decode(out, f"{output_name}.png", dst)
|
||||||
|
save_file(model.state_dict(), f"{output_name}.safetensors")
|
||||||
|
torch.save(optim.state_dict(), f"{output_name}.optim.pth")
|
||||||
|
|
||||||
|
class LatentDataset(Dataset):
|
||||||
|
class Shard:
|
||||||
|
def __init__(self, root, fname, res, src, dst):
|
||||||
|
self.fname = fname
|
||||||
|
self.src_path = f"{root}/{src}_{res}px/{fname}.npy"
|
||||||
|
self.dst_path = f"{root}/{dst}_{res}px/{fname}.npy"
|
||||||
|
|
||||||
|
def __init__(self, res, src, dst, root="latents"):
|
||||||
|
print("Loading latents from disk")
|
||||||
|
self.latents = []
|
||||||
|
for i in tqdm(os.listdir(f"{root}/{src}_{res}px")):
|
||||||
|
fname, ext = os.path.splitext(i)
|
||||||
|
assert ext == ".npy"
|
||||||
|
s = self.Shard(root, fname, res, src, dst)
|
||||||
|
if os.path.isfile(s.src_path) and os.path.isfile(s.dst_path):
|
||||||
|
self.latents.append(s)
|
||||||
|
|
||||||
|
def __len__(self):
|
||||||
|
return len(self.latents)
|
||||||
|
|
||||||
|
def __getitem__(self, index):
|
||||||
|
s = self.latents[index]
|
||||||
|
src = self.load_latent(s.src_path)
|
||||||
|
dst = self.load_latent(s.dst_path)
|
||||||
|
return (src, dst)
|
||||||
|
|
||||||
|
def load_latent(self, path):
|
||||||
|
lat = torch.from_numpy(np.load(path))
|
||||||
|
if lat.shape[0] == 1:
|
||||||
|
lat = torch.squeeze(lat, 0)
|
||||||
|
assert not torch.isnan(torch.sum(lat.float()))
|
||||||
|
return lat
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
args = parse_args()
|
args = parse_args()
|
||||||
target_dev = "cuda"
|
target_dev = "cuda"
|
||||||
latent_src = args.src
|
resolution = 768
|
||||||
latent_dst = args.dst
|
|
||||||
|
|
||||||
latents = load_latents(latent_src, latent_dst, target_dev)
|
dataset = LatentDataset(resolution, args.src, args.dst)
|
||||||
|
loader = DataLoader(
|
||||||
|
dataset,
|
||||||
|
batch_size=args.bs,
|
||||||
|
shuffle=True,
|
||||||
|
num_workers=0,
|
||||||
|
# num_workers=4,
|
||||||
|
# persistent_workers=True,
|
||||||
|
)
|
||||||
|
eval_src, eval_dst = get_eval_data(
|
||||||
|
dataset,
|
||||||
|
f"latents/test_{args.src}_{resolution}px.npy",
|
||||||
|
f"latents/test_{args.dst}_{resolution}px.npy",
|
||||||
|
target_dev,
|
||||||
|
)
|
||||||
|
|
||||||
if not os.path.isdir("models"): os.mkdir("models")
|
os.makedirs("models", exist_ok=True)
|
||||||
log = open(f"models/{latent_src}-to-{latent_dst}_interposer.csv", "w")
|
log = open(f"models/{args.src}-to-{args.dst}_interposer.csv", "w")
|
||||||
|
|
||||||
if os.path.isfile(f"test_{latent_src}.npy") and os.path.isfile(f"test_{latent_dst}.npy"):
|
|
||||||
ss_latent = torch.from_numpy(np.load(f"test_{latent_src}.npy")).to(target_dev)
|
|
||||||
st_latent = torch.from_numpy(np.load(f"test_{latent_dst}.npy")).to(target_dev)
|
|
||||||
else:
|
|
||||||
sample_latent = random.choice(latents)
|
|
||||||
ss_latent = sample_latent.src.to(target_dev)
|
|
||||||
st_latent = sample_latent.dst.to(target_dev)
|
|
||||||
|
|
||||||
model = Interposer()
|
model = Interposer()
|
||||||
|
|
||||||
|
criterion = torch.nn.L1Loss()
|
||||||
|
optimizer = torch.optim.AdamW(model.parameters(), lr=float(args.lr))
|
||||||
|
# import bitsandbytes as bnb
|
||||||
|
# optimizer = bnb.optim.AdamW8bit(model.parameters(), lr=float(args.lr))
|
||||||
|
|
||||||
|
scheduler = None
|
||||||
|
if args.cosine:
|
||||||
|
print("Using CosineAnnealingLR")
|
||||||
|
scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(
|
||||||
|
optimizer, T_max = int(args.steps/args.bs),
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
print("Using LinearLR")
|
||||||
|
scheduler = torch.optim.lr_scheduler.LinearLR(
|
||||||
|
optimizer,
|
||||||
|
start_factor = 0.1,
|
||||||
|
end_factor = 1.0,
|
||||||
|
total_iters = int(5000/args.bs),
|
||||||
|
)
|
||||||
|
|
||||||
if args.resume:
|
if args.resume:
|
||||||
model.load_state_dict(load_file(args.resume))
|
model.load_state_dict(load_file(args.resume))
|
||||||
model.to(target_dev)
|
model.to(target_dev)
|
||||||
|
optimizer.load_state_dict(torch.load(
|
||||||
|
f"{os.path.splitext(args.resume)[0]}.optim.pth"
|
||||||
|
))
|
||||||
|
else:
|
||||||
|
model.to(target_dev)
|
||||||
|
|
||||||
criterion = torch.nn.MSELoss(size_average=False)
|
progress = tqdm(total=args.steps)
|
||||||
optimizer = torch.optim.SGD(model.parameters(), lr=float(args.lr)/args.bs)
|
while progress.n < args.steps:
|
||||||
|
for src, dst in loader:
|
||||||
|
src = src.to(target_dev)
|
||||||
|
dst = dst.to(target_dev)
|
||||||
|
with torch.cuda.amp.autocast():
|
||||||
|
y_pred = model(src) # forward
|
||||||
|
loss = criterion(y_pred, dst) # loss
|
||||||
|
|
||||||
for t in tqdm(range(int(args.steps/args.bs)), unit_scale=args.bs):
|
# backward
|
||||||
step = t*args.bs
|
optimizer.zero_grad()
|
||||||
# input batch
|
loss.backward()
|
||||||
lts = [random.choice(latents) for _ in range(args.bs)]
|
optimizer.step()
|
||||||
src = torch.cat([x.src for x in lts],0)
|
scheduler.step()
|
||||||
dst = torch.cat([x.dst for x in lts],0)
|
|
||||||
|
|
||||||
y_pred = model(src) # forward
|
# eval/save
|
||||||
loss = criterion(y_pred, dst) # loss
|
progress.update(args.bs)
|
||||||
|
if progress.n % (1000 + 1000%args.bs) == 0:
|
||||||
|
eval_model(progress.n, model, criterion, scheduler, eval_src, eval_dst)
|
||||||
|
if progress.n % (args.save + args.save%args.bs) == 0:
|
||||||
|
save_model(progress.n, model, optimizer, eval_src, args.src, args.dst)
|
||||||
|
if progress.n >= args.steps:
|
||||||
|
break
|
||||||
|
progress.close()
|
||||||
|
|
||||||
# backward
|
|
||||||
optimizer.zero_grad()
|
|
||||||
loss.backward()
|
|
||||||
optimizer.step()
|
|
||||||
|
|
||||||
# print loss
|
|
||||||
if step%1000 == 0:
|
|
||||||
# test loss
|
|
||||||
with torch.no_grad():
|
|
||||||
t_pred = model(ss_latent)
|
|
||||||
t_loss = criterion(t_pred, st_latent)
|
|
||||||
tqdm.write(f"{step} - {loss.data.item()/args.bs:.2f}|{t_loss.data.item()/args.bs:.2f}")
|
|
||||||
log.write(f"{step},{loss.data.item()/args.bs:.2f},{t_loss.data.item()/args.bs:.2f}\n")
|
|
||||||
log.flush()
|
|
||||||
|
|
||||||
# sample/save
|
|
||||||
if step%args.save == 0:
|
|
||||||
out = model(ss_latent)
|
|
||||||
output_name = f"./models/{latent_src}-to-{latent_dst}_interposer_e{step/1000}k"
|
|
||||||
sample_decode(out, f"{output_name}.png", latent_dst)
|
|
||||||
save_file(model.state_dict(), f"{output_name}.safetensors")
|
|
||||||
# save final output
|
# save final output
|
||||||
output_name = f"./models/{latent_src}-to-{latent_dst}_interposer_e{args.steps/1000}k"
|
eval_model(progress.n, model, criterion, scheduler, eval_src, eval_dst)
|
||||||
sample_decode(out, f"{output_name}.png", "v1")
|
save_model(progress.n, model, optimizer, eval_src, args.src, args.dst)
|
||||||
save_file(model.state_dict(), f"{output_name}.safetensors")
|
|
||||||
log.close()
|
log.close()
|
||||||
|
|||||||
Reference in New Issue
Block a user