From 12bd221dbe816c255099ef45dea1679faa2f3889 Mon Sep 17 00:00:00 2001 From: aiXander Date: Thu, 14 Mar 2024 21:05:46 -0700 Subject: [PATCH] trying to fix trainer --- trainer/dataset_and_utils.py | 2 +- trainer/utils/rendering.py | 4 +++- trainer_pti.py | 6 +++--- 3 files changed, 7 insertions(+), 5 deletions(-) diff --git a/trainer/dataset_and_utils.py b/trainer/dataset_and_utils.py index 0cfe2d3..aeac08c 100755 --- a/trainer/dataset_and_utils.py +++ b/trainer/dataset_and_utils.py @@ -35,7 +35,7 @@ def plot_torch_hist(parameters, epoch, checkpoint_dir, name, bins=100, min_val=- plt.xlim(min_val, max_val) plt.xlabel('Weight Value') plt.ylabel('Count') - plt.title(f'Epoch {epoch} {name} Histogram (std = {np.std(all_params_cpu):.4f})') + plt.title(f'Epoch {epoch} {name} Histogram (std = {np.std(all_params_cpu):.4f}, min = {np.min(all_params_cpu):.2f}, max = {np.max(all_params_cpu):.2f})') plt.savefig(f"{checkpoint_dir}/{name}_histogram_{epoch:04d}.png") plt.close() diff --git a/trainer/utils/rendering.py b/trainer/utils/rendering.py index 336dab4..676972e 100644 --- a/trainer/utils/rendering.py +++ b/trainer/utils/rendering.py @@ -48,6 +48,8 @@ def make_validation_img_grid(img_folder): @torch.no_grad() def render_images(training_pipeline, render_size, lora_path, train_step, seed, is_lora, pretrained_model, lora_scale = 0.7, n_steps = 25, n_imgs = 4, device = "cuda:0"): + random.seed(seed) + with open(os.path.join(lora_path, "training_args.json"), "r") as f: training_args = json.load(f) concept_mode = training_args["concept_mode"] @@ -89,7 +91,7 @@ def render_images(training_pipeline, render_size, lora_path, train_step, seed, i pipeline.scheduler = EulerDiscreteScheduler.from_config(pipeline.scheduler.config) validation_prompts = [prepare_prompt_for_lora(prompt, lora_path) for prompt in validation_prompts_raw] - generator = torch.Generator(device=device).manual_seed(seed) + generator = torch.Generator(device=device).manual_seed(0) pipeline_args = { "negative_prompt": "nude, naked, poorly drawn face, ugly, tiling, out of frame, extra limbs, disfigured, deformed body, blurry, blurred, watermark, text, grainy, signature, cut off, draft", "num_inference_steps": n_steps, diff --git a/trainer_pti.py b/trainer_pti.py index 2b559e7..e18a67b 100755 --- a/trainer_pti.py +++ b/trainer_pti.py @@ -4,7 +4,7 @@ import os from io_utils import MODEL_DICT out_root_dir = "./lora_models" -run_name = "clipx_testing" +run_name = "clipx_ti_only" concept_mode = "style" output_dir = os.path.join(out_root_dir, run_name) @@ -46,7 +46,7 @@ config = (TrainerConfig)( textual_inversion_lr = 1e-3, textual_inversion_weight_decay = 3e-4, lora_weight_decay = 0.002, - prodigy_d_coef = 0.0, + prodigy_d_coef = 1.0, l1_penalty = 0.1, snr_gamma = 5.0, mixed_precision = "bf16", @@ -54,7 +54,7 @@ config = (TrainerConfig)( inserting_list_tokens = ["",""], is_lora = True, lora_rank = 12, - lora_alpha = 12 * 10, + lora_alpha = 12, hard_pivot = False, off_ratio_power = 0.1, args_dict = {},