remove objects
@@ -0,0 +1,87 @@
|
||||
#!/bin/bash
|
||||
|
||||
python main.py scripts/gridsearch_configs/faces/faces_000.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_001.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_002.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_003.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_004.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_005.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_006.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_007.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_008.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_009.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_010.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_011.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_012.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_013.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_014.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_015.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_016.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_017.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_018.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_019.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_020.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_021.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_022.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_023.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_024.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_025.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_026.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_027.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_028.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_029.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_030.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_031.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_032.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_033.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_034.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_035.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_036.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_037.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_038.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_039.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_040.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_041.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_042.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_043.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_044.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_045.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_046.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_047.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_048.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_049.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_050.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_051.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_052.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_053.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_054.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_055.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_056.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_057.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_058.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_059.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_060.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_061.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_062.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_063.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_064.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_065.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_066.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_067.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_068.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_069.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_070.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_071.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_072.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_073.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_074.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_075.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_076.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_077.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_078.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_079.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_080.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_081.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_082.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_083.json
|
||||
python main.py scripts/gridsearch_configs/faces/faces_084.json
|
||||
|
Before Width: | Height: | Size: 102 KiB |
|
Before Width: | Height: | Size: 531 KiB |
|
Before Width: | Height: | Size: 93 KiB |
|
Before Width: | Height: | Size: 228 KiB |
|
Before Width: | Height: | Size: 94 KiB |
|
Before Width: | Height: | Size: 102 KiB |
|
Before Width: | Height: | Size: 98 KiB |
|
Before Width: | Height: | Size: 120 KiB |
|
Before Width: | Height: | Size: 108 KiB |
@@ -1,236 +0,0 @@
|
||||
import argparse
|
||||
from trainer.inference import render_images_eval
|
||||
from trainer.utils.json_stuff import save_as_json
|
||||
from trainer.config import TrainingConfig
|
||||
from trainer.models import pretrained_models
|
||||
from trainer.utils.io import download
|
||||
import clip
|
||||
from PIL import Image
|
||||
import torch
|
||||
import numpy as np
|
||||
import os
|
||||
from creator_lora.models.resnet50 import ResNet50MLP
|
||||
|
||||
"""
|
||||
todos:
|
||||
- run eval on user-defined captions
|
||||
"""
|
||||
|
||||
aesthetic_model_checkpoint_filename = "aesthetic_score_best_model.pth"
|
||||
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
|
||||
def get_filenames_in_a_folder(folder: str):
|
||||
"""
|
||||
returns the list of paths to all the files in a given folder
|
||||
"""
|
||||
|
||||
if folder[-1] == '/':
|
||||
folder = folder[:-1]
|
||||
|
||||
files = os.listdir(folder)
|
||||
files = [f'{folder}/' + x for x in files]
|
||||
return files
|
||||
|
||||
def get_all_jpg_filenames(folder):
|
||||
all_filenames = get_filenames_in_a_folder(folder=folder)
|
||||
jpg_filenames = [filename for filename in all_filenames if filename.lower().endswith('.jpg')]
|
||||
assert len(jpg_filenames)>0, f"Expected to find at least 1 jpg file but got 0"
|
||||
return jpg_filenames
|
||||
|
||||
def filter_prompt(prompt, remove_this = "in the style of <s0><s1>,", replace_with = ""):
|
||||
assert remove_this in prompt, f"Expected '{remove_this}' to be present in the prompt: '{prompt}'"
|
||||
return prompt.replace(
|
||||
remove_this,
|
||||
replace_with
|
||||
)
|
||||
|
||||
def get_similarity_matrix(a, b, eps=1e-8):
|
||||
"""
|
||||
finds the cosine similarity matrix between each item of a w.r.t each item of b
|
||||
a and b are expected to be 2 dimensional
|
||||
added eps for numerical stability
|
||||
source: https://stackoverflow.com/a/58144658
|
||||
"""
|
||||
a_n, b_n = a.norm(dim=1)[:, None], b.norm(dim=1)[:, None]
|
||||
a_norm = a / torch.max(a_n, eps * torch.ones_like(a_n))
|
||||
b_norm = b / torch.max(b_n, eps * torch.ones_like(b_n))
|
||||
sim_mt = torch.mm(a_norm, b_norm.transpose(0, 1))
|
||||
return sim_mt
|
||||
|
||||
class Evaluation:
|
||||
def __init__(self, image_filenames: list):
|
||||
self.image_filenames = image_filenames
|
||||
self.image_features = None
|
||||
|
||||
def obtain_image_features(self):
|
||||
|
||||
if self.image_features is None:
|
||||
all_image_features = []
|
||||
model, preprocess = clip.load("ViT-B/32", device=device)
|
||||
|
||||
for f in self.image_filenames:
|
||||
image = preprocess(Image.open(f)).unsqueeze(0).to(device)
|
||||
with torch.no_grad():
|
||||
image_features = model.encode_image(image)
|
||||
all_image_features.append(image_features.float())
|
||||
|
||||
all_image_features = torch.cat(all_image_features, dim = 0)
|
||||
self.image_features = all_image_features
|
||||
|
||||
return self.image_features
|
||||
|
||||
def obtain_text_features(self, prompts: list, device):
|
||||
model, preprocess = clip.load("ViT-B/32", device=device)
|
||||
text = clip.tokenize(prompts).to(device)
|
||||
|
||||
with torch.no_grad():
|
||||
text_features = model.encode_text(text)
|
||||
return text_features
|
||||
|
||||
def training_image_alignment(self, device, training_image_filenames: list):
|
||||
generated_image_features = self.obtain_image_features()
|
||||
|
||||
training_image_features = []
|
||||
model, preprocess = clip.load("ViT-B/32", device=device)
|
||||
|
||||
for f in training_image_filenames:
|
||||
image = preprocess(Image.open(f)).unsqueeze(0).to(device)
|
||||
with torch.no_grad():
|
||||
image_features = model.encode_image(image)
|
||||
training_image_features.append(image_features.float())
|
||||
|
||||
training_image_features = torch.cat(training_image_features, dim = 0)
|
||||
return get_similarity_matrix(a=generated_image_features, b=training_image_features).mean().item()
|
||||
|
||||
|
||||
def image_text_alignment(self, device, prompts: list):
|
||||
|
||||
image_features = self.obtain_image_features().to(device)
|
||||
assert image_features.shape[0] == len(prompts), f'Expected len(prompts) ({len(prompts)}) to have the same number of prompts as the number of images provided: {image_features.shape}'
|
||||
text_features = self.obtain_text_features(prompts=prompts, device=device)
|
||||
cossim = torch.nn.functional.cosine_similarity(
|
||||
text_features, image_features, dim = -1
|
||||
).mean().item()
|
||||
return cossim
|
||||
|
||||
def clip_diversity(self, device: str):
|
||||
"""
|
||||
higher = more diverse
|
||||
"""
|
||||
all_image_features = self.obtain_image_features().to(device)
|
||||
|
||||
distances = 1 - get_similarity_matrix(all_image_features, all_image_features)
|
||||
assert distances.shape == (
|
||||
all_image_features.shape[0],
|
||||
all_image_features.shape[0]
|
||||
), f'Expected the shape of the distance matrix to be (num_images, num_images) i.e {(all_image_features.shape[0], all_image_features.shape[0])} but got: {distances.shape}'
|
||||
distances = distances.detach().cpu().numpy()
|
||||
# Get the upper triangle:
|
||||
upper_triangle = np.triu(distances, k=1).flatten()
|
||||
return upper_triangle.mean().item()
|
||||
|
||||
def aesthetic_score(self, device: str, checkpoint_path: str):
|
||||
# assert os.path.exists(checkpoint_path), f"invalid checkpoint_path: {checkpoint_path}"
|
||||
model = ResNet50MLP(
|
||||
model_path=checkpoint_path,
|
||||
device = device
|
||||
)
|
||||
|
||||
scores = []
|
||||
for f in self.image_filenames:
|
||||
score = model.predict_score(pil_image=Image.open(f))
|
||||
scores.append(score)
|
||||
|
||||
return sum(scores)/len(scores)
|
||||
|
||||
def parse_arguments():
|
||||
parser = argparse.ArgumentParser(description="Script for generating images based on prompts and computing similarities.")
|
||||
|
||||
|
||||
parser.add_argument("--config_filename", type=str, required=True, default = "sdxl", help="path to config json file")
|
||||
parser.add_argument("--checkpoint_folder", type=str, required=True,
|
||||
help="Path to folder containing the checkpoint. Usually a folder which is named like: .../checkpoint-500")
|
||||
parser.add_argument("--output_json", type=str, required=True,
|
||||
help="Path to json where we save result values")
|
||||
parser.add_argument("--output_folder", type=str, required=True,
|
||||
help="style or face")
|
||||
parser.add_argument("--training_images_folder", type=str, required=True,
|
||||
help="path to folder containing training image jpg files. Usually the `images_in` folder")
|
||||
args = parser.parse_args()
|
||||
|
||||
## validate args
|
||||
assert os.path.exists(args.checkpoint_folder), f"Invalid lora_path: {args.checkpoint_folder}"
|
||||
assert os.path.exists(args.config_filename), f"Invalid lora_path: {args.config_filename}"
|
||||
assert os.path.exists(args.training_images_folder), f"Invalid training_images_folder: {args.training_images_folder}"
|
||||
return args
|
||||
|
||||
args = parse_arguments()
|
||||
|
||||
os.system(f"mkdir -p {args.output_folder}")
|
||||
if not os.path.exists(aesthetic_model_checkpoint_filename):
|
||||
download(
|
||||
url="https://edenartlab-lfs.s3.amazonaws.com/models/aesthetic_score_best_model.pth",
|
||||
folder="./",
|
||||
filepath=None
|
||||
)
|
||||
|
||||
config = TrainingConfig.from_json(args.config_filename)
|
||||
|
||||
image_filenames, prompts = render_images_eval(
|
||||
output_folder=args.output_folder,
|
||||
concept_mode=config.concept_mode,
|
||||
render_size=(1024,1024),
|
||||
checkpoint_folder=args.checkpoint_folder,
|
||||
pretrained_model=pretrained_models[config.sd_model_version],
|
||||
seed=0,
|
||||
is_lora = config.is_lora,
|
||||
trigger_text='TOK' if config.concept_mode != "style" else ", in the style of TOK"
|
||||
)
|
||||
|
||||
|
||||
print(f"Eval prompts:")
|
||||
for i, p in enumerate(prompts):
|
||||
print(f"{i}:{p}")
|
||||
|
||||
eval = Evaluation(image_filenames=image_filenames)
|
||||
clip_diversity = eval.clip_diversity(device=device)
|
||||
|
||||
|
||||
aesthetic_score = eval.aesthetic_score(device=device, checkpoint_path=aesthetic_model_checkpoint_filename)
|
||||
image_text_alignment = eval.image_text_alignment(device=device, prompts=prompts)
|
||||
training_image_alignment = eval.training_image_alignment(
|
||||
device=device,
|
||||
training_image_filenames=get_all_jpg_filenames(folder=args.training_images_folder)
|
||||
)
|
||||
|
||||
result = {
|
||||
"sd_model_version": config.sd_model_version,
|
||||
"checkpoint_folder": os.path.abspath(args.checkpoint_folder),
|
||||
"concept_mode": config.concept_mode,
|
||||
"output_folder": args.output_folder,
|
||||
"training_images_folder":args.training_images_folder,
|
||||
"scores": {
|
||||
"clip_diversity": clip_diversity,
|
||||
"aesthetic_score": aesthetic_score,
|
||||
"image_text_alignment": image_text_alignment,
|
||||
"training_image_alignment": training_image_alignment
|
||||
}
|
||||
}
|
||||
|
||||
save_as_json(
|
||||
dictionary_or_list=result,
|
||||
filename=args.output_json
|
||||
)
|
||||
print(f"Eval complete. Saved results here: {args.output_json}")
|
||||
|
||||
"""
|
||||
Example command:
|
||||
|
||||
python3 evaluate.py \
|
||||
--output_folder eval_images \
|
||||
--checkpoint_folder lora_models/clipx--17_05-20-54-sdxl_style_dora_512_1.0_blip/checkpoints/checkpoint-0 \
|
||||
--output_json eval_results_style.json \
|
||||
--config_filename lora_models/clipx--17_05-20-54-sdxl_style_dora_512_1.0_blip/checkpoints/checkpoint-0/training_args.json \
|
||||
--training_images_folder lora_models/clipx--17_05-20-54-sdxl_style_dora_512_1.0_blip/images_in
|
||||
"""
|
||||
@@ -1,142 +0,0 @@
|
||||
import os
|
||||
import json
|
||||
import matplotlib.pyplot as plt
|
||||
import seaborn as sns
|
||||
from collections import defaultdict
|
||||
import numpy as np
|
||||
import matplotlib.pyplot as plt
|
||||
import seaborn as sns
|
||||
from sklearn.linear_model import LinearRegression
|
||||
from sklearn.metrics import r2_score
|
||||
|
||||
|
||||
# Define paths
|
||||
render_dir = "/home/rednax/SSD2TB/Xander_Tools/sd15_face_sweep/lora_models"
|
||||
config_dir = "/home/rednax/SSD2TB/Xander_Tools/sd15_face_sweep/xander_adiff_lora"
|
||||
|
||||
ignore_threshold_relative = 0.0 # ignore any datapoint with a score below this threshold
|
||||
|
||||
filters = {
|
||||
"resolution": 512
|
||||
}
|
||||
|
||||
output_dir = f"gridsearch_configs/results/{os.path.basename(config_dir)}"
|
||||
output_suffix = f"{os.path.basename(render_dir)}"
|
||||
|
||||
# Initialize a dictionary to hold parameter values and associated scores
|
||||
parameters = defaultdict(lambda: defaultdict(list))
|
||||
|
||||
# Step 1: Loop over each experiment subdirectory
|
||||
for i, exp_subdir in enumerate(sorted(os.listdir(render_dir))):
|
||||
exp_path = os.path.join(render_dir, exp_subdir)
|
||||
checkpoints_path = os.path.join(exp_path, "checkpoints")
|
||||
|
||||
# Step 2: Get the score by counting the number of .jpg files in the checkpoints subdir
|
||||
if os.path.isdir(checkpoints_path):
|
||||
score = sum(1 for _ in os.listdir(checkpoints_path) if _.endswith('.jpg'))
|
||||
|
||||
# Match the experiment folder with its corresponding JSON file
|
||||
json_file_name = exp_subdir.split('--')[0] + ".json"
|
||||
json_file_name = json_file_name.replace('__','_')
|
||||
json_path = os.path.join(config_dir, json_file_name)
|
||||
|
||||
# Step 3: Load the corresponding .json file
|
||||
if os.path.isfile(json_path):
|
||||
with open(json_path, 'r') as file:
|
||||
config = json.load(file)
|
||||
|
||||
# Filter out experiments that do not match the filters
|
||||
if not all(config[key] == value for key, value in filters.items()):
|
||||
continue
|
||||
|
||||
# Step 4: Append all key/value pairs to the total experiment dictionary
|
||||
for key, value in config.items():
|
||||
parameters[key]['values'].append(value)
|
||||
parameters[key]['scores'].append(score)
|
||||
else:
|
||||
print(f"Could not find JSON file for experiment {exp_subdir}")
|
||||
|
||||
|
||||
# Print the parameters['output_dir'] with the highest scores (there are usually multiple ties):
|
||||
max_score = max(parameters['output_dir']['scores'])
|
||||
best_output_dirs = [output_dir for output_dir, score in zip(parameters['output_dir']['values'], parameters['output_dir']['scores']) if score == max_score]
|
||||
for best_output_dir in best_output_dirs:
|
||||
print(f"Best output_dir: {best_output_dir} with score {max_score}")
|
||||
|
||||
import numpy as np
|
||||
import matplotlib.pyplot as plt
|
||||
import seaborn as sns
|
||||
from sklearn.linear_model import LinearRegression
|
||||
from sklearn.metrics import r2_score
|
||||
from sklearn.preprocessing import LabelEncoder
|
||||
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
print(f"Saving results to {output_dir}...")
|
||||
|
||||
def plot_parameters(parameters):
|
||||
for param, data in parameters.items():
|
||||
values = np.array(data['values'])
|
||||
scores = np.array(data['scores'])
|
||||
|
||||
# filter based on the ignore_threshold:
|
||||
ignore_threshold = ignore_threshold_relative * np.max(scores)
|
||||
mask = scores > ignore_threshold
|
||||
values = values[mask]
|
||||
scores = scores[mask]
|
||||
|
||||
noise_strength_values = 0.02
|
||||
noise_strength_scores = 0.02
|
||||
|
||||
# Initialize variables for original categorical labels
|
||||
original_labels = None
|
||||
|
||||
# Determine if values are numeric
|
||||
if values.dtype.kind in 'bifc': # Numeric types
|
||||
# Add noise directly to values
|
||||
jittered_values = values + np.random.normal(0, noise_strength_values * (np.max(values) - np.min(values)), values.shape)
|
||||
else:
|
||||
# Encode string values to integers for plotting
|
||||
encoder = LabelEncoder()
|
||||
original_labels = values.copy()
|
||||
values = encoder.fit_transform(values)
|
||||
jittered_values = values + np.random.normal(0, 0.1, values.shape)
|
||||
|
||||
# Skip plotting if there is only one unique value for the parameter
|
||||
if len(np.unique(values)) <= 1:
|
||||
continue
|
||||
|
||||
# Fit a linear regression model to the encoded values if categorical
|
||||
model = LinearRegression()
|
||||
values_reshaped = values.reshape(-1, 1) # Reshape for sklearn
|
||||
model.fit(values_reshaped, scores)
|
||||
predicted_scores = model.predict(values_reshaped)
|
||||
|
||||
# add some jitter to the scores:
|
||||
jittered_scores = scores + np.random.normal(0, noise_strength_scores * np.max(scores), scores.shape)
|
||||
|
||||
# Calculate R² value
|
||||
r_squared = r2_score(scores, predicted_scores)
|
||||
|
||||
# Plot data points
|
||||
sns.scatterplot(x=jittered_values, y=jittered_scores, alpha=0.6)
|
||||
|
||||
# Plot trendline
|
||||
sns.lineplot(x=np.sort(values), y=predicted_scores[np.argsort(values)], color='red', label=f'R²={r_squared:.2f}')
|
||||
|
||||
# Set plot title and labels
|
||||
plt.title(f'Influence of {param} on the score')
|
||||
if original_labels is not None:
|
||||
# Set x-axis labels to the original categorical labels
|
||||
unique_values = np.unique(values)
|
||||
plt.xticks(ticks=unique_values, labels=encoder.inverse_transform(unique_values), rotation=45, ha='right')
|
||||
else:
|
||||
plt.xlabel(param)
|
||||
plt.ylabel('Score')
|
||||
plt.legend()
|
||||
|
||||
# Save and close the plot
|
||||
plt.savefig(f'{output_dir}/res_{param}_{output_suffix}.png')
|
||||
plt.close()
|
||||
|
||||
# Call the updated function with your parameters dictionary
|
||||
plot_parameters(parameters)
|
||||
@@ -1,146 +0,0 @@
|
||||
import os
|
||||
import json
|
||||
import matplotlib.pyplot as plt
|
||||
from collections import defaultdict
|
||||
|
||||
# Step 1: Import necessary libraries (done above)
|
||||
|
||||
# Step 2: Define a function to count JPG files in a directory
|
||||
def count_jpg_files(directory):
|
||||
return len([f for f in os.listdir(directory) if f.lower().endswith('.jpg')])
|
||||
|
||||
# Step 3: Define a function to find and load the training_args.json file
|
||||
def load_training_args(directory):
|
||||
for root, dirs, files in os.walk(directory):
|
||||
if 'training_args.json' in files:
|
||||
with open(os.path.join(root, 'training_args.json'), 'r') as f:
|
||||
return json.load(f)
|
||||
return None
|
||||
|
||||
# Step 4: Traverse the directory structure and collect data
|
||||
def collect_data(root_dir):
|
||||
data = []
|
||||
for root, dirs, files in os.walk(root_dir):
|
||||
if 'checkpoints' in dirs:
|
||||
checkpoints_dir = os.path.join(root, 'checkpoints')
|
||||
score = count_jpg_files(checkpoints_dir)
|
||||
training_args = load_training_args(root)
|
||||
if training_args:
|
||||
data.append((training_args, score))
|
||||
else:
|
||||
print(f"Warning: No training_args.json found in {root}")
|
||||
print(f"Collected data from {len(data)} runs")
|
||||
return data
|
||||
|
||||
# Step 5: Process the collected data to identify varying hyperparameters
|
||||
def identify_varying_hyperparams(data):
|
||||
all_params = set().union(*[set(args.keys()) for args, _ in data])
|
||||
varying_params = {}
|
||||
|
||||
def make_hashable(val):
|
||||
if isinstance(val, dict):
|
||||
return tuple(sorted((k, make_hashable(v)) for k, v in val.items()))
|
||||
elif isinstance(val, list):
|
||||
return tuple(make_hashable(v) for v in val)
|
||||
elif isinstance(val, set):
|
||||
return frozenset(make_hashable(v) for v in val)
|
||||
return val
|
||||
|
||||
for param in all_params:
|
||||
try:
|
||||
values = set(make_hashable(args.get(param)) for args, _ in data if param in args)
|
||||
if len(values) > 1:
|
||||
varying_params[param] = values
|
||||
print(f"---> Parameter '{param}' varies across runs")
|
||||
except TypeError as e:
|
||||
print(f"Warning: Could not hash values for parameter '{param}'. Error: {e}")
|
||||
print(f"Values: {[args.get(param) for args, _ in data if param in args]}")
|
||||
|
||||
return varying_params
|
||||
|
||||
|
||||
# Step 6: Create visual plots for each varying hyperparameter
|
||||
import numpy as np
|
||||
from scipy import stats
|
||||
import matplotlib.pyplot as plt
|
||||
from collections import defaultdict
|
||||
|
||||
def create_plots(data, varying_params):
|
||||
for param, values in varying_params.items():
|
||||
plt.figure(figsize=(12, 8))
|
||||
param_data = defaultdict(list)
|
||||
|
||||
for args, score in data:
|
||||
if param in args:
|
||||
value = args[param]
|
||||
value_str = str(value)
|
||||
param_data[value_str].append(score)
|
||||
|
||||
values_list = sorted(param_data.keys())
|
||||
all_x = []
|
||||
all_y = []
|
||||
|
||||
for i, value_str in enumerate(values_list):
|
||||
scores = param_data[value_str]
|
||||
jittered_x = np.random.normal(i, 0.1, size=len(scores))
|
||||
plt.scatter(jittered_x, scores, alpha=0.6, label=value_str)
|
||||
all_x.extend([i] * len(scores))
|
||||
all_y.extend(scores)
|
||||
|
||||
# Calculate trendline
|
||||
x = np.array(all_x)
|
||||
y = np.array(all_y)
|
||||
z = np.polyfit(x, y, 1)
|
||||
p = np.poly1d(z)
|
||||
|
||||
# Calculate R-squared
|
||||
r_squared = 1 - (sum((y - p(x))**2) / ((len(y) - 1) * np.var(y, ddof=1)))
|
||||
|
||||
# Plot trendline
|
||||
plt.plot(x, p(x), "r--", alpha=0.8,
|
||||
label=f'Trendline: y={z[0]:.2f}x+{z[1]:.2f}\nR²: {r_squared:.4f}')
|
||||
|
||||
plt.xlabel(param)
|
||||
plt.ylabel('Score')
|
||||
plt.title(f'Effect of {param} on Score')
|
||||
|
||||
# Adjust x-axis labels
|
||||
if len(values_list) > 10:
|
||||
plt.xticks(range(0, len(values_list), len(values_list)//10),
|
||||
[values_list[i] for i in range(0, len(values_list), len(values_list)//10)],
|
||||
rotation=45, ha='right')
|
||||
else:
|
||||
plt.xticks(range(len(values_list)), values_list, rotation=45, ha='right')
|
||||
|
||||
# Adjust legend
|
||||
if len(values_list) > 10:
|
||||
plt.legend(title="Legend", bbox_to_anchor=(1.05, 1), loc='upper left')
|
||||
else:
|
||||
plt.legend(bbox_to_anchor=(1.05, 1), loc='upper left')
|
||||
|
||||
plt.tight_layout()
|
||||
|
||||
# Save figure with error handling
|
||||
try:
|
||||
plt.savefig(f'{param}_vs_score.png', dpi=200, bbox_inches='tight')
|
||||
except ValueError:
|
||||
print(f"Warning: Failed to save image for {param}. skipping..")
|
||||
|
||||
|
||||
plt.close()
|
||||
|
||||
print("Plots have been saved as PNG files in the current directory.")
|
||||
|
||||
if __name__ == "__main__":
|
||||
root_dir = "/home/rednax/SSD2TB/Github_repos/diffusion_trainer/lora_models/PLANTOID"
|
||||
|
||||
# Collect data
|
||||
data = collect_data(root_dir)
|
||||
|
||||
# Identify varying hyperparameters
|
||||
varying_params = identify_varying_hyperparams(data)
|
||||
|
||||
# Create plots
|
||||
create_plots(data, varying_params)
|
||||
|
||||
print("Plots have been saved as PNG files in the current directory.")
|
||||