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:
@@ -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()
|
||||||
|
|||||||
Reference in New Issue
Block a user