| from collections import Counter |
| import json |
| from sklearn.model_selection import train_test_split |
| from sklearn.model_selection import StratifiedKFold |
| from sklearn.model_selection import KFold |
| from modules.data_cleaner import get_valid_instances |
|
|
| def generate_splits(fsl_labels_json: str, output_filepath: str, seed = 42): |
| final_dict = {} |
| input_features = [] |
| output_labels =[] |
|
|
| labels_dict = get_valid_instances(fsl_labels_json) |
|
|
| for x in labels_dict.keys(): |
| for y in labels_dict[x]["instances"]: |
| input_features.append(y["filepath"]) |
| output_labels.append(x) |
|
|
| train_input_features, test_input_features, train_output_labels, test_output_labels = train_test_split( |
| input_features, output_labels, test_size=0.2, random_state=seed, stratify=output_labels |
| ) |
|
|
| test_list = [] |
| test_list = generate_dict(test_input_features, test_output_labels) |
| final_dict.update({"test":test_list}) |
|
|
| skf = StratifiedKFold(n_splits=5, shuffle=True, random_state=seed) |
|
|
| with open(output_filepath, "w") as jsonfile: |
| train_list = [] |
| validate_list = [] |
|
|
| for i, (train_index, validate_index) in enumerate( |
| skf.split(train_input_features, train_output_labels) |
| ): |
| fold_input_features = [train_input_features[j] for j in train_index] |
| fold_output_labels = [train_output_labels[j] for j in train_index] |
|
|
| train_list = generate_dict(fold_input_features, fold_output_labels) |
|
|
| fold_input_features = [train_input_features[j] for j in validate_index] |
| fold_output_labels = [train_output_labels[j] for j in validate_index] |
|
|
| validate_list = generate_dict(fold_input_features, fold_output_labels) |
|
|
| fold_dict = ({"train": train_list, "validate": validate_list}) |
| final_dict.update({"fold_" + str(i):fold_dict}) |
|
|
|
|
|
|
| jsonfile.write(json.dumps(final_dict, indent=2)) |
|
|
|
|
| def generate_dict(input_features, output_labels): |
| output_list = [] |
|
|
| for x, y in zip(input_features, output_labels): |
| temp_dict = {"filepath":x, "class": y} |
| output_list.append(temp_dict) |
|
|
| return output_list |
|
|
| def continuous_data_splitter(json_path, output_filepath, n_splits=5): |
| print(f"Loading metadata from {json_path}...") |
| with open(json_path, 'r') as f: |
| data = json.load(f) |
| |
| sample_keys = list(data["samples"].keys()) |
| |
| |
| kf = KFold(n_splits=n_splits, shuffle=True, random_state=42) |
| |
| splits_output = {} |
| |
| print(f"Generating {n_splits}-fold cross-validation splits...") |
| |
| fold_idx = 1 |
| for train_index, val_index in kf.split(sample_keys): |
| |
| |
| train_keys = [sample_keys[i] for i in train_index] |
| val_keys = [sample_keys[i] for i in val_index] |
| |
| splits_output[f"fold_{fold_idx}"] = { |
| "train": train_keys, |
| "validation": val_keys |
| } |
| |
| print(f" Fold {fold_idx}: {len(train_keys)} Train, {len(val_keys)} Validation") |
| fold_idx += 1 |
| |
| with open(output_filepath, "w") as jsonfile: |
| jsonfile.write(json.dumps(splits_output, indent=2)) |
|
|