fsl-express / modules /dataset_splitter.py
lasofeli's picture
Upload folder using huggingface_hub
bc971c7 verified
Raw
History Blame Contribute Delete
3.24 kB
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())
# Initialize KFold
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):
# Convert indices back to 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))