CoolFace
Apppublic

baulab/Erasing-Concepts-In-Diffusion

sourceHugging Facemitupdated 3y agoView on Hugging Face
49likes
train.py91 linesDownload Raw Back to root
1from StableDiffuser import StableDiffuser2from finetuning import FineTunedModel3import torch4from tqdm import tqdm5 6def train(prompt, modules, freeze_modules, iterations, negative_guidance, lr, save_path):7  8    nsteps = 509 10    diffuser = StableDiffuser(scheduler='DDIM').to('cuda')11    diffuser.train()12 13    finetuner = FineTunedModel(diffuser, modules, frozen_modules=freeze_modules)14 15    optimizer = torch.optim.Adam(finetuner.parameters(), lr=lr)16    criteria = torch.nn.MSELoss()17 18    pbar = tqdm(range(iterations))19 20    with torch.no_grad():21 22        neutral_text_embeddings = diffuser.get_text_embeddings([''],n_imgs=1)23        positive_text_embeddings = diffuser.get_text_embeddings([prompt],n_imgs=1)24 25    del diffuser.vae26    del diffuser.text_encoder27    del diffuser.tokenizer28 29    torch.cuda.empty_cache()30 31    for i in pbar:32        33        with torch.no_grad():34 35            diffuser.set_scheduler_timesteps(nsteps)36 37            optimizer.zero_grad()38 39            iteration = torch.randint(1, nsteps - 1, (1,)).item()40 41            latents = diffuser.get_initial_latents(1, 512, 1)42 43            with finetuner:44 45                latents_steps, _ = diffuser.diffusion(46                    latents,47                    positive_text_embeddings,48                    start_iteration=0,49                    end_iteration=iteration,50                    guidance_scale=3, 51                    show_progress=False52                )53 54            diffuser.set_scheduler_timesteps(1000)55 56            iteration = int(iteration / nsteps * 1000)57            58            positive_latents = diffuser.predict_noise(iteration, latents_steps[0], positive_text_embeddings, guidance_scale=1)59            neutral_latents = diffuser.predict_noise(iteration, latents_steps[0], neutral_text_embeddings, guidance_scale=1)60 61        with finetuner:62            negative_latents = diffuser.predict_noise(iteration, latents_steps[0], positive_text_embeddings, guidance_scale=1)63 64        positive_latents.requires_grad = False65        neutral_latents.requires_grad = False66 67        loss = criteria(negative_latents, neutral_latents - (negative_guidance*(positive_latents - neutral_latents))) #loss = criteria(e_n, e_0) works the best try 5000 epochs68        69        loss.backward()70        optimizer.step()71 72    torch.save(finetuner.state_dict(), save_path)73 74    del diffuser, loss, optimizer, finetuner, negative_latents, neutral_latents, positive_latents, latents_steps, latents75 76    torch.cuda.empty_cache()77if __name__ == '__main__':78 79    import argparse80 81    parser = argparse.ArgumentParser()82 83    parser.add_argument('--prompt', required=True)84    parser.add_argument('--modules', required=True)85    parser.add_argument('--freeze_modules', nargs='+', required=True)86    parser.add_argument('--save_path', required=True)87    parser.add_argument('--iterations', type=int, required=True)88    parser.add_argument('--lr', type=float, required=True)89    parser.add_argument('--negative_guidance', type=float, required=True)90 91    train(**vars(parser.parse_args()))