update grid_img display

This commit is contained in:
xander
2024-08-07 22:21:45 +02:00
parent c31b122686
commit c527fc690d
2 changed files with 17 additions and 8 deletions
+16 -7
View File
@@ -97,8 +97,6 @@ class Eden_LoRa_trainer:
config, output_save_dir = e.value # Capture the return value
break
validation_grid_img_path = os.path.join(output_save_dir, "validation_grid.jpg")
attributes = {}
attributes['grid_prompts'] = config.training_attributes["validation_prompts"]
attributes['job_time_seconds'] = config.job_time
@@ -116,11 +114,22 @@ class Eden_LoRa_trainer:
else:
lora_path = path
# Load the grid image:
grid_image = Image.open(validation_grid_img_path)
grid_image = np.array(grid_image).astype(np.float32) / 255.0
grid_image = torch.from_numpy(grid_image)[None,]
# Load the grid images:
grid_images = []
grid_dir = os.path.dirname(output_save_dir)
for f in os.listdir(grid_dir):
if "validation_grid" in f:
grid_image = Image.open(os.path.join(grid_dir, f))
grid_image = np.array(grid_image).astype(np.float32) / 255.0
grid_image = torch.from_numpy(grid_image)
grid_images.append(grid_image)
grid_images = torch.stack(grid_images)
# Make sure that grid_images always has 4 dimensions:
if len(grid_images.shape) == 3:
grid_images = grid_images.unsqueeze(0)
final_msg = f"LoRa trained in {config.job_time/60:.1f} minutes. Files saved at {output_save_dir}"
return (grid_image, lora_path, embedding_path, final_msg)
return (grid_images, lora_path, embedding_path, final_msg)
+1 -1
View File
@@ -1,5 +1,5 @@
{
"name": "gene_sdxl",
"name": "xander_sdxl",
"sd_model_version": "sdxl",
"lora_training_urls": "https://storage.googleapis.com/public-assets-xander/A_workbox/lora_training_sets/xander.zip",
"concept_mode": "face",