Add batch size

This is probably not how it's supposed to be done but it stops the horrible coil whine at least.
This commit is contained in:
City
2023-07-31 21:41:18 +02:00
parent 994bb1a04b
commit 8d619afdf8
+22 -14
View File
@@ -14,6 +14,8 @@ from vae import get_vae
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('--lr', default="1e-8", 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")
@@ -21,6 +23,10 @@ def parse_args():
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")
try:
float(args.lr)
except:
parser.error("--lr must be a valid float eg. 0.001 or 1e-3")
return args return args
class Latent: class Latent:
@@ -63,8 +69,6 @@ def sample_decode(latent, filename, version):
if __name__ == "__main__": if __name__ == "__main__":
args = parse_args() args = parse_args()
target_dev = "cuda" target_dev = "cuda"
target_steps = args.steps
save_every_n = args.save
latent_src = args.src latent_src = args.src
latent_dst = args.dst latent_dst = args.dst
@@ -84,14 +88,17 @@ if __name__ == "__main__":
model.to(target_dev) model.to(target_dev)
criterion = torch.nn.MSELoss(size_average=False) criterion = torch.nn.MSELoss(size_average=False)
optimizer = torch.optim.SGD(model.parameters(), lr=1e-8) optimizer = torch.optim.SGD(model.parameters(), lr=float(args.lr)/args.bs)
for t in tqdm(range(target_steps)): for t in tqdm(range(int(args.steps/args.bs)), unit_scale=args.bs):
# io = latents[t%len(latents)] step = t*args.bs
io = random.choice(latents) # input batch
lts = [random.choice(latents) for _ in range(args.bs)]
src = torch.cat([x.src for x in lts],0)
dst = torch.cat([x.dst for x in lts],0)
y_pred = model(io.src) # forward y_pred = model(src) # forward
loss = criterion(y_pred, io.dst) # loss loss = criterion(y_pred, dst) # loss
# backward # backward
optimizer.zero_grad() optimizer.zero_grad()
@@ -99,18 +106,19 @@ if __name__ == "__main__":
optimizer.step() optimizer.step()
# print loss # print loss
if t%1000 == 0: if step%1000 == 0:
tqdm.write(f"{t} - {loss.data.item():.2f}") tqdm.write(f"{step} - {loss.data.item()/args.bs:.2f}")
log.write(f"{t},{loss.data.item():.2f}\n") log.write(f"{step},{loss.data.item()/args.bs:.2f}\n")
log.flush()
# sample/save # sample/save
if t%save_every_n == 0: if step%args.save == 0:
out = model(sample_latent) out = model(sample_latent)
output_name = f"./models/{latent_src}-to-{latent_dst}_interposer_e{t/1000}k" output_name = f"./models/{latent_src}-to-{latent_dst}_interposer_e{step/1000}k"
sample_decode(out, f"{output_name}.png", latent_dst) sample_decode(out, f"{output_name}.png", latent_dst)
save_file(model.state_dict(), f"{output_name}.safetensors") 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{target_steps/1000}k" output_name = f"./models/{latent_src}-to-{latent_dst}_interposer_e{args.steps/1000}k"
sample_decode(out, f"{output_name}.png", "v1") sample_decode(out, f"{output_name}.png", "v1")
save_file(model.state_dict(), f"{output_name}.safetensors") save_file(model.state_dict(), f"{output_name}.safetensors")
log.close() log.close()