From 50203abb35a9ef63b242dbb3ae4e511f8b2c0dc1 Mon Sep 17 00:00:00 2001 From: Mohamad Zamini <32536264+mzamini92@users.noreply.github.com> Date: Wed, 19 Jul 2023 21:58:12 -0600 Subject: [PATCH] Update generate_configs.py some improvements on generate_configs file. Instead of directly comparing the dataset names in get_data_handler_config, we can use constants to make the code more readable and maintainable. Rather than passing strings for model types like "sklearn" we used enums to represent different model types. Improved the serialization technique. --- .../generate_configs.py | 75 ++++++++++++------- 1 file changed, 48 insertions(+), 27 deletions(-) diff --git a/examples/sklearn_logclassification_rw/generate_configs.py b/examples/sklearn_logclassification_rw/generate_configs.py index 37238b2..747b946 100644 --- a/examples/sklearn_logclassification_rw/generate_configs.py +++ b/examples/sklearn_logclassification_rw/generate_configs.py @@ -1,11 +1,23 @@ import os - import joblib +from enum import Enum from sklearn.linear_model import SGDClassifier import examples.datahandlers as datahandlers +class Dataset(Enum): + ADULT = "adult" + COMPAS = "compas" + GERMAN = "german" + CUSTOM_DATASET = "custom_dataset" + + +class ModelType(Enum): + SKLEARN = "sklearn" + # Add other model types here if needed + + def get_fusion_config(): fusion = {"name": "IterAvgFusionHandler", "path": "ibmfl.aggregator.fusion.iter_avg_fusion_handler"} return fusion @@ -19,45 +31,54 @@ def get_local_training_config(configs_folder=None): return local_training_handler -def get_hyperparams(model): - hyperparams = {"global": {"rounds": 3, "termination_accuracy": 0.9}, "local": {"training": {"max_iter": 2}}} - +def get_hyperparams(model_type: ModelType): + hyperparams = { + "global": {"rounds": 3, "termination_accuracy": 0.9}, + "local": {"training": {"max_iter": 2}} + } return hyperparams -def get_data_handler_config(party_id, dataset, folder_data, is_agg=False, model="sklearn"): - SUPPORTED_DATASETS = ["adult", "compas", "german", "custom_dataset"] - - if dataset in SUPPORTED_DATASETS: - if dataset == "adult": - dataset = "adult_sklearn" - elif dataset == "compas": - dataset = "compas_sklearn" - elif dataset == "german": - dataset = "german_sklearn" - data = datahandlers.get_datahandler_config(dataset, folder_data, party_id, is_agg) - else: - raise Exception("The dataset {} is a wrong combination for fusion/model".format(dataset)) - return data +def get_data_handler_config(party_id, dataset: Dataset, folder_data, is_agg=False): + SUPPORTED_DATASETS = { + Dataset.ADULT: "adult_sklearn", + Dataset.COMPAS: "compas_sklearn", + Dataset.GERMAN: "german_sklearn", + } + if dataset not in SUPPORTED_DATASETS: + raise Exception("The dataset {} is not supported.".format(dataset.value)) -def get_model_config(folder_configs, dataset, is_agg=False, party_id=0, model="sklearn"): - if is_agg: - return None + data = datahandlers.get_datahandler_config(SUPPORTED_DATASETS[dataset], folder_data, party_id, is_agg) + return data - model = SGDClassifier(loss="log", penalty="l2") +def save_model(model, folder_configs): if not os.path.exists(folder_configs): os.makedirs(folder_configs) fname = os.path.join(folder_configs, "model_architecture.pickle") - with open(fname, "wb") as f: joblib.dump(model, f) - # Generate model spec: - spec = {"model_definition": fname} + return fname - model = {"name": "SklearnSGDFLModel", "path": "ibmfl.model.sklearn_SGD_linear_fl_model", "spec": spec} - return model +def get_model_config(folder_configs, dataset: Dataset, is_agg=False, party_id=0, model_type: ModelType = ModelType.SKLEARN): + if is_agg: + return None + + if model_type == ModelType.SKLEARN: + model = SGDClassifier(loss="log", penalty="l2") + else: + raise ValueError("Model type {} is not supported.".format(model_type.value)) + + model_spec = {"model_definition": save_model(model, folder_configs)} + + model_config = { + "name": "SklearnSGDFLModel", + "path": "ibmfl.model.sklearn_SGD_linear_fl_model", + "spec": model_spec, + } + + return model_config