From fadf4d8f22cf903833d09a267a3e3e91208d9fa6 Mon Sep 17 00:00:00 2001 From: Fennel1 <627235787@qq.com> Date: Wed, 29 Apr 2026 18:37:50 +0800 Subject: [PATCH 1/9] add dataset & model --- apps/backend/app/api/v1/endpoints/agent.py | 37 +- .../app/services/agent/capabilities.py | 115 +--- .../app/services/agent/experiment_service.py | 3 + apps/backend/app/services/agent/graph.py | 20 +- apps/backend/app/services/agent/planning.py | 9 +- .../app/services/agent/prompts/__init__.py | 5 +- .../agent/prompts/plan_instructions.txt | 8 +- .../services/distributed/config_service.py | 2 +- .../app/services/simulation/compatibility.py | 58 ++ .../app/services/simulation/job_service.py | 3 + apps/backend/runners/core_runtime.py | 69 ++- .../tests/api/v1/endpoints/test_agent.py | 5 +- .../tests/services/agent/test_planning.py | 24 +- .../tests/services/test_job_service.py | 78 +++ .../components/AgentExperimentStudio.tsx | 41 +- .../agent/components/AgentPlanPreview.tsx | 40 +- apps/frontend/src/features/agent/schema.ts | 58 +- .../src/features/agent/useAgentController.ts | 4 +- .../src/features/config/compatibility.ts | 199 +++++++ .../features/config/render-schema-section.tsx | 71 ++- config_schema.yaml | 102 +++- configs/simulation_config.yaml | 14 +- libs/fl_core/compression/__init__.py | 19 + libs/fl_core/compression/sparsification.py | 169 +++++- libs/fl_core/data/data_loader.py | 456 +++++++++++---- libs/fl_core/federated/client_manager.py | 28 +- libs/fl_core/models/__init__.py | 13 +- libs/fl_core/models/basic.py | 122 ++++ libs/fl_core/models/model_manager.py | 316 ++++++++--- libs/fl_core/models/resnet.py | 8 +- libs/fl_core/privacy/__init__.py | 9 + libs/fl_core/privacy/differential_privacy.py | 54 ++ libs/fl_core/privacy/secure_aggregation.py | 54 ++ libs/fl_core/simulation_registry.py | 528 ++++++++++++++++++ libs/fl_core/utils/config.py | 8 +- 35 files changed, 2338 insertions(+), 411 deletions(-) create mode 100644 apps/backend/app/services/simulation/compatibility.py create mode 100644 apps/frontend/src/features/config/compatibility.ts create mode 100644 libs/fl_core/models/basic.py create mode 100644 libs/fl_core/privacy/differential_privacy.py create mode 100644 libs/fl_core/privacy/secure_aggregation.py create mode 100644 libs/fl_core/simulation_registry.py diff --git a/apps/backend/app/api/v1/endpoints/agent.py b/apps/backend/app/api/v1/endpoints/agent.py index 29f6f1b..db3130f 100644 --- a/apps/backend/app/api/v1/endpoints/agent.py +++ b/apps/backend/app/api/v1/endpoints/agent.py @@ -27,7 +27,13 @@ from app.services.agent import AgentExperimentRunService, AgentExperimentService, AgentOptimizationHistoryService, AgentRuntimeService, AgentState from app.services.agent.graph import FederatedAgentGraphBuilder from app.services.agent.objectives import AgentOptimizationObjective, resolve_objective -from app.services.agent.planning import build_schema_prompt_context, dumps_for_prompt +from app.services.agent.planning import ( + build_initial_config, + build_schema_prompt_context, + deep_merge_config, + dumps_for_prompt, + lock_structured_constraints, +) from app.services.llm import LLMRegistry, LLMService from app.schemas.agent import AgentPlanPreviewRequest, AgentPlanPreviewResponse, ExperimentPlanPreview, AgentPlanReviseRequest import uuid @@ -675,13 +681,38 @@ async def revise_plan_preview( except Exception as e: raise HTTPException(status_code=500, detail=f"Failed to parse LLM response: {str(e)}") - snapshot["draft_experiments"] = updated_experiments + experiment_service = AgentExperimentService(session) + schema = experiment_service.get_config_schema() + constrained_base = experiment_service.normalize_simulation_config( + deep_merge_config(build_initial_config(schema), config_constraints) + ) + normalized_experiments = [] + for idx, exp in enumerate(updated_experiments): + if not isinstance(exp, dict): + continue + raw_config = exp.get("config_patch", exp.get("config", {})) + if not isinstance(raw_config, dict): + raw_config = {} + merged = lock_structured_constraints( + deep_merge_config(constrained_base, raw_config), + config_constraints, + ) + normalized_experiments.append( + { + **exp, + "name": exp.get("name", f"exp-{idx + 1}"), + "plan_summary": exp.get("plan_summary", snapshot.get("goal", "")), + "config_patch": experiment_service.normalize_simulation_config(merged), + } + ) + + snapshot["draft_experiments"] = normalized_experiments await history_service.update_job_snapshot(task_id=job.task_id, snapshot=snapshot) return AgentPlanPreviewResponse( optimization_job_id=job.id, goal=snapshot.get("goal", ""), - experiments=[ExperimentPlanPreview.model_validate(exp) for exp in updated_experiments], + experiments=[ExperimentPlanPreview.model_validate(exp) for exp in normalized_experiments], system_mode=snapshot.get("system_mode", "simulation"), config_constraints=config_constraints, ) diff --git a/apps/backend/app/services/agent/capabilities.py b/apps/backend/app/services/agent/capabilities.py index 4b973e2..f3fae88 100644 --- a/apps/backend/app/services/agent/capabilities.py +++ b/apps/backend/app/services/agent/capabilities.py @@ -1,118 +1,23 @@ """ -Helpers for inspecting the current federated learning stack capabilities. +Helpers for exposing executable simulation capabilities to the Agent. -Used by the Agent to ground its proposals in the actually supported -datasets, models, and aggregations. +The Agent should plan only combinations that the simulation runtime can +execute. The source of truth lives in libs/fl_core/simulation_registry.py; +this module is a thin backend-side adapter. """ from __future__ import annotations -import ast -import re from pathlib import Path -from typing import Any, Dict, List +from typing import Any, Dict - -def _read_text(path: Path) -> str: - return path.read_text(encoding="utf-8", errors="replace") - - -def _extract_list_literal(text: str) -> List[Any]: - """ - Extract a Python list literal from a string like: [ "a", "b" ]. - Returns [] on failure. - """ - try: - value = ast.literal_eval(text) - if isinstance(value, list): - return value - except Exception: - pass - return [] - - -def _extract_supported_datasets(project_root: Path) -> List[str]: - path = project_root / "libs" / "fl_core" / "data" / "data_loader.py" - if not path.exists(): - return [] - text = _read_text(path) - match = re.search(r"supported_datasets\s*=\s*(\[[^\]]*\])", text, flags=re.DOTALL) - if not match: - return [] - items = _extract_list_literal(match.group(1)) - out: List[str] = [] - for item in items: - if isinstance(item, str): - out.append(item) - return out - - -def _extract_model_registry_keys(project_root: Path) -> List[str]: - path = project_root / "libs" / "fl_core" / "models" / "model_manager.py" - if not path.exists(): - return [] - text = _read_text(path) - match = re.search(r"model_registry\s*=\s*\{(.*?)\}\s*\n", text, flags=re.DOTALL) - if not match: - return [] - block = match.group(1) - keys = re.findall(r"['\"]([a-zA-Z0-9_]+)['\"]\s*:", block) - seen: set[str] = set() - out: List[str] = [] - for key in keys: - lower = key.lower() - if lower not in seen: - seen.add(lower) - out.append(lower) - return out - - -def _extract_aggregation_strategies(project_root: Path) -> List[str]: - """ - Returns aggregation strategies from federated/aggregation.py without imports. - """ - path = project_root / "libs" / "fl_core" / "federated" / "aggregation.py" - if not path.exists(): - return [] - text = _read_text(path) - - match = re.search(r"_strategies\s*=\s*\{(.*?)\}", text, flags=re.DOTALL) - if not match: - return [] - aggregations = re.findall(r"['\"]([a-zA-Z0-9_]+)['\"]\s*:", match.group(1)) - - seen: set[str] = set() - out: List[str] = [] - for item in aggregations: - lower = item.lower() - if lower not in seen: - seen.add(lower) - out.append(lower) - return out +from app.services.simulation.compatibility import runtime_capabilities def get_platform_capabilities(project_root: Path | None = None) -> Dict[str, Any]: """ - Inspect the current codebase and return the supported configuration options - for datasets, models, and aggregations. + Return supported simulation datasets, models, aggregations, and compatibility + rules. The project_root argument is kept for compatibility with older tests. """ - root = project_root or Path(__file__).resolve().parents[4] - - datasets = _extract_supported_datasets(root) - models = _extract_model_registry_keys(root) - aggregations = _extract_aggregation_strategies(root) - - metrics = { - "global_results": ["rounds", "global_loss", "global_accuracy"], - "client_results": ["train_loss", "train_acc", "test_loss", "test_acc"], - } - - return { - "datasets": datasets, - "distributions": ["iid", "non_iid"], - "models": models, - "aggregations": aggregations, - "metrics": metrics, - } - - + _ = project_root + return runtime_capabilities() diff --git a/apps/backend/app/services/agent/experiment_service.py b/apps/backend/app/services/agent/experiment_service.py index a6bbd8d..6dc2040 100644 --- a/apps/backend/app/services/agent/experiment_service.py +++ b/apps/backend/app/services/agent/experiment_service.py @@ -32,6 +32,7 @@ AgentExperimentStatus, ) from app.repositories.agent import AgentExperimentRepository +from app.services.simulation.compatibility import canonicalize_runtime_config, validate_runtime_config_or_raise from app.core.logger import get_logger @@ -235,7 +236,9 @@ def _normalize_simulation_config(cls, config: dict[str, Any]) -> dict[str, Any]: system["mode"] = "simulation" system["node_role"] = "server" + canonicalize_runtime_config(merged) cls._validate_config_node(merged, schema, path_prefix="") + validate_runtime_config_or_raise(merged) return merged @classmethod diff --git a/apps/backend/app/services/agent/graph.py b/apps/backend/app/services/agent/graph.py index f9d305f..7663a43 100644 --- a/apps/backend/app/services/agent/graph.py +++ b/apps/backend/app/services/agent/graph.py @@ -31,6 +31,7 @@ build_schema_prompt_context, collect_disabled_option_errors, deep_merge_config, + lock_structured_constraints, select_llm_model, ) from .prompts import build_plan_prompt, build_plan_system_instructions @@ -97,8 +98,12 @@ async def _node_parse(self, state: AgentState) -> AgentState: raw_config = exp.get("config_patch", exp.get("config", {})) if not isinstance(raw_config, dict): raw_config = {} + merged = lock_structured_constraints( + deep_merge_config(constrained_base, raw_config), + state.config_constraints, + ) normalized = experiment_service.normalize_simulation_config( - deep_merge_config(constrained_base, raw_config) + merged ) disabled_errors = collect_disabled_option_errors(normalized, schema) if disabled_errors: @@ -165,8 +170,12 @@ async def _node_parse(self, state: AgentState) -> AgentState: patch_config = exp.get("config", {}) if not isinstance(patch_config, dict): patch_config = {} - # Merge with constrained base config and normalize. - merged = deep_merge_config(constrained_base, patch_config) + # Merge with constrained base config, then re-apply + # structured controls so prompt drift cannot override them. + merged = lock_structured_constraints( + deep_merge_config(constrained_base, patch_config), + state.config_constraints, + ) try: normalized = experiment_service.normalize_simulation_config(merged) disabled_errors = collect_disabled_option_errors(normalized, schema) @@ -232,7 +241,10 @@ async def _node_run_sequential(self, state: AgentState) -> AgentState: state.iteration = idx + 1 # -- Build config -- - merged = deep_merge_config(base_config, plan.config_patch) + merged = lock_structured_constraints( + deep_merge_config(base_config, plan.config_patch), + state.config_constraints, + ) config = experiment_service.normalize_simulation_config(merged) disabled_errors = collect_disabled_option_errors(config, schema) if disabled_errors: diff --git a/apps/backend/app/services/agent/planning.py b/apps/backend/app/services/agent/planning.py index c66570c..73016e0 100644 --- a/apps/backend/app/services/agent/planning.py +++ b/apps/backend/app/services/agent/planning.py @@ -49,7 +49,7 @@ def build_initial_config(schema: dict[str, Any]) -> dict[str, Any]: if not schema: return { "dataset": {"name": "CIFAR-10", "distribution": "non_iid", "alpha": 0.5}, - "model": {"name": "CNN"}, + "model": {"name": "Auto"}, "federated": { "num_clients": 3, "num_rounds": 10, @@ -88,6 +88,13 @@ def deep_merge_config(base: dict[str, Any], override: dict[str, Any] | None) -> return merged +def lock_structured_constraints(config: dict[str, Any], constraints: dict[str, Any] | None) -> dict[str, Any]: + """Apply user-selected structured constraints as final locked values.""" + if not isinstance(constraints, dict) or not constraints: + return copy.deepcopy(config) + return deep_merge_config(config, constraints) + + def _option_ui(definition: dict[str, Any]) -> dict[str, Any]: ui = definition.get("ui") if not isinstance(ui, dict): diff --git a/apps/backend/app/services/agent/prompts/__init__.py b/apps/backend/app/services/agent/prompts/__init__.py index 73a3276..4b1dd98 100644 --- a/apps/backend/app/services/agent/prompts/__init__.py +++ b/apps/backend/app/services/agent/prompts/__init__.py @@ -48,7 +48,10 @@ def build_plan_prompt( f"Default base config after applying user constraints:\n{dumps_for_prompt(base_config)}\n\n" f"User-selected structured constraints:\n{dumps_for_prompt(constraints)}\n\n" "Parse the request into experiment configurations.\n" - "Every experiment must inherit the structured constraints unless the user's request explicitly changes them.\n" + "User-selected structured constraints are locked controls and have higher priority than natural-language text.\n" + "Every experiment must inherit the structured constraints. If text conflicts with them, keep the structured constraints.\n" "Use executable_options for select fields. Do not use disabled_options.\n" + "Respect dataset_model_compatibility. Use model.name=\"Auto\" when the user changes dataset without naming a model.\n" + "Respect privacy/compression compatibility rules from Runtime capabilities.\n" "Return ONLY a JSON object with plan_summary and experiments list, no commentary." ) diff --git a/apps/backend/app/services/agent/prompts/plan_instructions.txt b/apps/backend/app/services/agent/prompts/plan_instructions.txt index 316aaa1..e062590 100644 --- a/apps/backend/app/services/agent/prompts/plan_instructions.txt +++ b/apps/backend/app/services/agent/prompts/plan_instructions.txt @@ -5,9 +5,13 @@ Rules: - Treat config_schema.yaml as the source of truth for fields, defaults, selectable options, and disabled future options. - If the user mentions multiple parameter values, generate one experiment per value. - If the user mentions multiple dimensions, generate the Cartesian product. -- If the user does not specify a field, inherit the default base config and structured constraints from the prompt. +- User-selected structured constraints are locked controls. Every experiment must inherit them, even if the natural-language text or an example mentions a different value. +- If the user does not specify a field, inherit the default base config from the prompt. - Each experiment's "config" should contain the actual parameter values that make that experiment distinct. - Do not use any option listed under disabled_options. +- Respect Runtime capabilities, especially dataset_model_compatibility and privacy/compression/aggregation compatibility. +- Never combine CKKS with compression or secure aggregation. Secure aggregation is only valid with fedavg/simple_avg in simulation. Differential privacy can be combined with supported compression/encryption. +- If the user changes dataset but does not ask for a model, use model.name="Auto" or the dataset's default_model. - `federated.seed` controls random seed for full reproducibility. Do not vary it unless the user asks for repeats, variance, or specific seeds. - Return ONLY a JSON object with: - "plan_summary": brief description of what will be tested @@ -15,7 +19,7 @@ Rules: Example: user says "test alpha=0.1, 0.3, 0.5": { - "plan_summary": "Compare non-IID alpha values 0.1, 0.3, 0.5 on CIFAR-10", + "plan_summary": "Compare non-IID alpha values 0.1, 0.3, 0.5 on the selected dataset", "experiments": [ {"name": "alpha_0.1", "config": {"dataset": {"alpha": 0.1}}}, {"name": "alpha_0.3", "config": {"dataset": {"alpha": 0.3}}}, diff --git a/apps/backend/app/services/distributed/config_service.py b/apps/backend/app/services/distributed/config_service.py index cd8e14e..b6e70e4 100644 --- a/apps/backend/app/services/distributed/config_service.py +++ b/apps/backend/app/services/distributed/config_service.py @@ -117,7 +117,7 @@ def _init_defaults_from_schema(cls, schema: dict[str, Any]) -> dict[str, Any]: """ output: dict[str, Any] = {} for key, raw_definition in schema.items(): - if key in {"role", "depends_on", "hidden"}: + if key in {"role", "depends_on", "hidden", "ui"}: continue if not isinstance(raw_definition, dict): continue diff --git a/apps/backend/app/services/simulation/compatibility.py b/apps/backend/app/services/simulation/compatibility.py new file mode 100644 index 0000000..0066c63 --- /dev/null +++ b/apps/backend/app/services/simulation/compatibility.py @@ -0,0 +1,58 @@ +from __future__ import annotations + +import sys +from pathlib import Path +from typing import Any + +from app.core import exceptions + + +def _project_root() -> Path: + return Path(__file__).resolve().parents[5] + + +def _ensure_fl_core_path() -> None: + libs_path = _project_root() / "libs" + if str(libs_path) not in sys.path: + sys.path.insert(0, str(libs_path)) + + +def _registry(): + _ensure_fl_core_path() + from fl_core import simulation_registry + + return simulation_registry + + +def canonicalize_runtime_config(config: dict[str, Any]) -> None: + registry = _registry() + + dataset = config.get("dataset") + if not isinstance(dataset, dict): + return + raw_dataset_name = dataset.get("name") + if isinstance(raw_dataset_name, str): + dataset["name"] = registry.canonical_dataset_name(raw_dataset_name) + + spec = registry.get_dataset_spec(dataset.get("name")) + + model = config.get("model") + if not isinstance(model, dict): + return + raw_model_name = model.get("name") + if isinstance(raw_model_name, str) and raw_model_name.strip().lower() != "auto": + model["name"] = registry.canonical_model_name(raw_model_name) + if spec is not None: + model["input_shape"] = list(spec.input_shape) + model["num_classes"] = spec.num_classes + + +def validate_runtime_config_or_raise(config: dict[str, Any]) -> None: + registry = _registry() + errors = registry.validate_training_combination(config) + if errors: + raise exceptions.BadRequestError("; ".join(errors)) + + +def runtime_capabilities() -> dict[str, Any]: + return _registry().capabilities_payload() diff --git a/apps/backend/app/services/simulation/job_service.py b/apps/backend/app/services/simulation/job_service.py index 59b4386..140ec57 100644 --- a/apps/backend/app/services/simulation/job_service.py +++ b/apps/backend/app/services/simulation/job_service.py @@ -15,6 +15,7 @@ from app.core import exceptions from app.models.simulation import SimulationJob, SimulationJobStatus, SimulationRunStatus from app.repositories.simulation import SimulationJobRepository, SimulationRunRepository +from app.services.simulation.compatibility import canonicalize_runtime_config, validate_runtime_config_or_raise class SimulationJobService: @@ -230,7 +231,9 @@ def _normalize_simulation_config(cls, config: dict[str, Any]) -> dict[str, Any]: system["mode"] = "simulation" system["node_role"] = "server" + canonicalize_runtime_config(merged) cls._validate_config_node(merged, schema, path_prefix="") + validate_runtime_config_or_raise(merged) return merged @classmethod diff --git a/apps/backend/runners/core_runtime.py b/apps/backend/runners/core_runtime.py index 48fce22..732eac3 100644 --- a/apps/backend/runners/core_runtime.py +++ b/apps/backend/runners/core_runtime.py @@ -22,11 +22,14 @@ from fl_core.data.data_loader import DataManager from fl_core.data.data_splitter import DataSplitter, FederatedDataManager from fl_core.models.model_manager import ModelManager +from fl_core.simulation_registry import resolve_model_name_for_dataset from fl_core.federated.server import FederatedServer from fl_core.federated.client import FederatedClient from fl_core.federated.client_manager import ClientManager -from fl_core.compression.sparsification import GlobalTopKSparsifier +from fl_core.compression.sparsification import CompressionFactory +from fl_core.privacy.differential_privacy import DifferentialPrivacyManager from fl_core.privacy.encryption import CKKSManager +from fl_core.privacy.secure_aggregation import SecureAggregationMasker from fl_core.federated.communication import GRPCServer, GRPCWorker, GRPCClientProxy WebApp = None @@ -96,14 +99,36 @@ def initialize_system(self) -> bool: self.sparsifier = None if sparsify_conf.get('enable', False): + method = sparsify_conf.get('method', 'global_topk') ratio = sparsify_conf.get('ratio', 0.5) - self.sparsifier = GlobalTopKSparsifier(ratio=ratio) + threshold = sparsify_conf.get('threshold', 1e-3) + self.sparsifier = CompressionFactory.create(method=method, ratio=ratio, threshold=threshold) print(f"✓ 通信压缩: 稀疏化已启用 (Ratio: {ratio})") privacy_config = self.config_manager.config.get('privacy', {}) ckks_conf = privacy_config.get('homomorphic_encryption', {}) - + dp_conf = privacy_config.get('differential_privacy', {}) + secure_agg_conf = privacy_config.get('secure_aggregation', {}) + + self.dp_manager = None + if dp_conf.get('enable', False): + self.dp_manager = DifferentialPrivacyManager( + clipping_norm=dp_conf.get('clipping_norm', 1.0), + noise_multiplier=dp_conf.get('noise_multiplier', 0.0), + ) + print( + "Differential privacy enabled: " + f"clip={self.dp_manager.clipping_norm}, noise={self.dp_manager.noise_multiplier}" + ) + + self.secure_aggregation = None + if secure_agg_conf.get('enable', False): + self.secure_aggregation = SecureAggregationMasker( + mask_std=secure_agg_conf.get('mask_std', 1.0), + ) + print(f"Secure aggregation masking enabled: std={self.secure_aggregation.mask_std}") + self.ckks_manager = None if ckks_conf.get('enable', False): # 暂不支持同时开启 @@ -155,7 +180,10 @@ def _get_experiment_info(self) -> Dict[str, Any]: "dataset_name": dataset_cfg.get('name', 'Unknown'), "distribution": dataset_cfg.get('distribution', 'iid'), "alpha": dataset_cfg.get('alpha', 'N/A'), - "model_name": model_cfg.get('name', 'Unknown'), + "model_name": resolve_model_name_for_dataset( + model_cfg.get('name'), + dataset_cfg.get('name', 'CIFAR-10'), + ), "num_classes": model_cfg.get('num_classes', 10) }, "federated": { @@ -169,13 +197,23 @@ def _get_experiment_info(self) -> Dict[str, Any]: "security": { "encryption": { "enabled": privacy_cfg.get('homomorphic_encryption', {}).get('enable', False), - "type": "CKKS", + "type": privacy_cfg.get('homomorphic_encryption', {}).get('method', 'ckks'), "poly_modulus_degree": privacy_cfg.get('homomorphic_encryption', {}).get('poly_modulus_degree', 0) }, + "differential_privacy": { + "enabled": privacy_cfg.get('differential_privacy', {}).get('enable', False), + "clipping_norm": privacy_cfg.get('differential_privacy', {}).get('clipping_norm', 1.0), + "noise_multiplier": privacy_cfg.get('differential_privacy', {}).get('noise_multiplier', 0.0), + }, + "secure_aggregation": { + "enabled": privacy_cfg.get('secure_aggregation', {}).get('enable', False), + "mask_std": privacy_cfg.get('secure_aggregation', {}).get('mask_std', 1.0), + }, "compression": { "enabled": compression_cfg.get('sparsification', {}).get('enable', False), - "type": "Global Top-K", - "ratio": compression_cfg.get('sparsification', {}).get('ratio', 0.0) + "type": compression_cfg.get('sparsification', {}).get('method', 'global_topk'), + "ratio": compression_cfg.get('sparsification', {}).get('ratio', 0.0), + "threshold": compression_cfg.get('sparsification', {}).get('threshold', 1e-3), } } } @@ -248,7 +286,9 @@ def run_distributed_server(self): self.client_manager = ClientManager( clients=[], sparsifier=self.sparsifier, - ckks_manager=self.ckks_manager + ckks_manager=self.ckks_manager, + dp_manager=self.dp_manager, + secure_aggregation=self.secure_aggregation, ) self.logger.log_info(f"等待 {federated_config.get('num_clients')} 个远程客户端连接...") @@ -422,11 +462,16 @@ def setup_models(self) -> bool: self.model_manager = ModelManager(model_config) dataset_info = self.data_manager.get_dataset_info() + model_name = resolve_model_name_for_dataset( + model_config.get('name'), + dataset_info['name'], + ) global_model = self.model_manager.create_model( - model_name=model_config.get('name'), + model_name=model_name, input_shape=dataset_info['input_shape'], - num_classes=dataset_info['num_classes'] + num_classes=dataset_info['num_classes'], + dataset_info=dataset_info, ) self.global_model = global_model @@ -479,6 +524,8 @@ def setup_federated_components(self) -> bool: clients=clients, sparsifier=self.sparsifier, ckks_manager=self.ckks_manager, + dp_manager=self.dp_manager, + secure_aggregation=self.secure_aggregation, selection_strategy="random" ) # Parallel client training races the global torch RNG inside @@ -524,7 +571,7 @@ def run_federated_training(self) -> bool: self.client_manager.broadcast_model_to_clients(all_clients, global_params) training_results = self.client_manager.train_clients( - selected_clients=all_clients, + selected_clients=selected_clients, epochs=local_epochs, learning_rate=learning_rate, round_num=round_num diff --git a/apps/backend/tests/api/v1/endpoints/test_agent.py b/apps/backend/tests/api/v1/endpoints/test_agent.py index 44d5b2b..cf62080 100644 --- a/apps/backend/tests/api/v1/endpoints/test_agent.py +++ b/apps/backend/tests/api/v1/endpoints/test_agent.py @@ -22,7 +22,10 @@ def test_agent_config_schema_endpoint_returns_ui_metadata(client): assert response.status_code == 200 payload = response.json() - assert payload["model"]["name"]["options"] == ["CNN", "LeNet", "ResNet"] + assert "Auto" in payload["model"]["name"]["options"] + assert "FedAvgCNN" in payload["model"]["name"]["options"] + assert "TextDNN" in payload["model"]["name"]["options"] + assert "CharLSTM" in payload["model"]["name"]["options"] aggregation_ui = payload["federated"]["aggregation"]["ui"] assert aggregation_ui["featured"] is True assert aggregation_ui["options"]["fedprox"]["disabled"] is True diff --git a/apps/backend/tests/services/agent/test_planning.py b/apps/backend/tests/services/agent/test_planning.py index 0fa66db..eb67d9c 100644 --- a/apps/backend/tests/services/agent/test_planning.py +++ b/apps/backend/tests/services/agent/test_planning.py @@ -3,6 +3,8 @@ build_initial_config, build_schema_prompt_context, collect_disabled_option_errors, + deep_merge_config, + lock_structured_constraints, select_llm_model, ) @@ -44,11 +46,31 @@ def test_build_initial_config_uses_schema_defaults(): def test_build_initial_config_hardcoded_defaults_on_empty_schema(): cfg = build_initial_config({}) assert cfg["dataset"]["name"] == "CIFAR-10" - assert cfg["model"]["name"] == "CNN" + assert cfg["model"]["name"] == "Auto" assert cfg["federated"]["num_clients"] == 3 assert cfg["federated"]["num_rounds"] == 10 +def test_lock_structured_constraints_override_llm_patch_values(): + constrained_base = { + "dataset": {"name": "FEMNIST", "alpha": 0.5}, + "model": {"name": "LeNet"}, + "federated": {"aggregation": "fedavg"}, + } + llm_patch = { + "dataset": {"name": "CIFAR-10", "alpha": 0.1}, + "model": {"name": "Auto"}, + } + constraints = {"dataset": {"name": "FEMNIST"}, "model": {"name": "LeNet"}} + + merged = deep_merge_config(constrained_base, llm_patch) + locked = lock_structured_constraints(merged, constraints) + + assert locked["dataset"]["name"] == "FEMNIST" + assert locked["dataset"]["alpha"] == 0.1 + assert locked["model"]["name"] == "LeNet" + + def test_schema_prompt_context_marks_disabled_options(): schema = { "federated": { diff --git a/apps/backend/tests/services/test_job_service.py b/apps/backend/tests/services/test_job_service.py index 9bd3220..5feec15 100644 --- a/apps/backend/tests/services/test_job_service.py +++ b/apps/backend/tests/services/test_job_service.py @@ -25,3 +25,81 @@ def test_normalize_simulation_config_keeps_dataset_and_federated_client_counts_a assert normalized["federated"]["num_clients"] == 5 assert normalized["dataset"]["num_clients"] == 5 + + +def test_normalize_simulation_config_allows_auto_model_for_dataset_switch(): + normalized = SimulationJobService.normalize_simulation_config( + { + "dataset": {"name": "MNIST"}, + "model": {"name": "Auto"}, + } + ) + + assert normalized["dataset"]["name"] == "MNIST" + assert normalized["model"]["name"] == "Auto" + assert normalized["model"]["input_shape"] == [1, 28, 28] + assert normalized["model"]["num_classes"] == 10 + + +def test_normalize_simulation_config_rejects_incompatible_model(): + with pytest.raises(exceptions.BadRequestError, match="not compatible"): + SimulationJobService.normalize_simulation_config( + { + "dataset": {"name": "AG News"}, + "model": {"name": "ResNet18"}, + } + ) + + +def test_normalize_simulation_config_rejects_ckks_with_sparsification(): + with pytest.raises(exceptions.BadRequestError, match="not compatible"): + SimulationJobService.normalize_simulation_config( + { + "privacy": {"homomorphic_encryption": {"enable": True}}, + "compression": {"sparsification": {"enable": True}}, + } + ) + + +def test_normalize_simulation_config_accepts_dp_with_compression(): + normalized = SimulationJobService.normalize_simulation_config( + { + "privacy": { + "differential_privacy": { + "enable": True, + "clipping_norm": 1.0, + "noise_multiplier": 0.1, + } + }, + "compression": { + "sparsification": { + "enable": True, + "method": "random_k", + "ratio": 0.25, + } + }, + } + ) + + assert normalized["privacy"]["differential_privacy"]["enable"] is True + assert normalized["compression"]["sparsification"]["method"] == "random_k" + + +def test_normalize_simulation_config_rejects_secure_aggregation_with_weighted_avg(): + with pytest.raises(exceptions.BadRequestError, match="equal-weight"): + SimulationJobService.normalize_simulation_config( + { + "federated": {"aggregation": "weighted_avg"}, + "privacy": {"secure_aggregation": {"enable": True}}, + } + ) + + +def test_normalize_simulation_config_rejects_secure_aggregation_with_compression(): + with pytest.raises(exceptions.BadRequestError, match="not compatible"): + SimulationJobService.normalize_simulation_config( + { + "privacy": {"secure_aggregation": {"enable": True}}, + "compression": {"sparsification": {"enable": True}}, + } + ) diff --git a/apps/frontend/src/features/agent/components/AgentExperimentStudio.tsx b/apps/frontend/src/features/agent/components/AgentExperimentStudio.tsx index ef8e760..7a356cd 100644 --- a/apps/frontend/src/features/agent/components/AgentExperimentStudio.tsx +++ b/apps/frontend/src/features/agent/components/AgentExperimentStudio.tsx @@ -12,17 +12,22 @@ import { Textarea } from "../../../components/ui/textarea"; import type { AgentPageProps } from "../../../pages/types"; import { getValueByPath } from "../../simulation/utils"; import { + booleanDisabledForConfig, + booleanDisableReasonForConfig, buildAgentConfig, checkDependency, coerceFieldValue, collectAgentSchemaFields, + fieldCompatibilityHint, formatFieldValue, - optionDisabled, + optionDisableReasonForConfig, + optionDisabledForConfig, optionLabel, optionMeta, selectedConstraintChips, type AgentSchemaField, } from "../schema"; +import { isModelCompatibleWithDataset } from "../../config/compatibility"; export function AgentExperimentStudio(props: AgentPageProps) { const { @@ -106,9 +111,19 @@ export function AgentExperimentStudio(props: AgentPageProps) { {visibleFields.map((field) => ( setConfigConstraint(field.path, coerceFieldValue(field, value))} + onChange={(value) => { + const coerced = coerceFieldValue(field, value); + setConfigConstraint(field.path, coerced); + if (field.path === "dataset.name") { + const currentModel = getValueByPath(effectiveConfig, "model.name"); + if (!isModelCompatibleWithDataset(currentModel, String(coerced))) { + setConfigConstraint("model.name", "Auto"); + } + } + }} /> ))} @@ -167,10 +182,12 @@ function parseOptionalNumber(value: unknown): number | undefined { } function ConstraintField({ + config, field, value, onChange, }: { + config: Record; field: AgentSchemaField; value: unknown; onChange: (value: unknown) => void; @@ -178,6 +195,9 @@ function ConstraintField({ const hint = field.definition.ui && typeof field.definition.ui.prompt_hint === "string" ? field.definition.ui.prompt_hint : null; + const compatibilityHint = fieldCompatibilityHint(field.path, config); + const boolDisabled = booleanDisabledForConfig(field.path, value, config); + const boolReason = booleanDisableReasonForConfig(field.path, value, config); return (
@@ -193,16 +213,21 @@ function ConstraintField({ {field.options.map((option) => { const meta = optionMeta(field.definition, option); + const disabled = optionDisabledForConfig(field.definition, option, field.path, config); + const reason = optionDisableReasonForConfig(field.definition, option, field.path, config); return ( - - {optionLabel(field.definition, option)} - {typeof meta.badge === "string" && {meta.badge}} + + + {optionLabel(field.definition, option)} + {typeof meta.badge === "string" && {meta.badge}} + + {reason && {reason}} ); @@ -223,12 +248,14 @@ function ConstraintField({ {field.type === "bool" && (
{formatFieldValue(value)} - +
)} {field.type !== "select" && field.type !== "number" && field.type !== "bool" && ( onChange(event.target.value)} /> )} + {compatibilityHint &&

{compatibilityHint}

} + {boolReason &&

{boolReason}

} {hint &&

{hint}

}
); diff --git a/apps/frontend/src/features/agent/components/AgentPlanPreview.tsx b/apps/frontend/src/features/agent/components/AgentPlanPreview.tsx index 10dc599..0f9785f 100644 --- a/apps/frontend/src/features/agent/components/AgentPlanPreview.tsx +++ b/apps/frontend/src/features/agent/components/AgentPlanPreview.tsx @@ -15,18 +15,23 @@ import { Textarea } from "../../../components/ui/textarea"; import type { AgentPageProps } from "../../../pages/types"; import { getValueByPath, isRecord } from "../../simulation/utils"; import { + booleanDisabledForConfig, + booleanDisableReasonForConfig, buildAgentConfig, checkDependency, coerceFieldValue, collectAgentSchemaFields, + fieldCompatibilityHint, formatFieldValue, - optionDisabled, + optionDisableReasonForConfig, + optionDisabledForConfig, optionLabel, optionMeta, setConfigValue, validateDisabledOptions, type AgentSchemaField, } from "../schema"; +import { isModelCompatibleWithDataset } from "../../config/compatibility"; type LocalExperiment = { name: string; @@ -91,7 +96,14 @@ export function AgentPlanPreview(props: AgentPageProps) { setLocalExperiments((current) => current.map((exp, idx) => { if (idx !== index) return exp; - const nextConfig = setConfigValue(exp.config_patch ?? {}, path, value); + let nextConfig = setConfigValue(exp.config_patch ?? {}, path, value); + if (path === "dataset.name") { + const fullConfig = buildAgentConfig(configSchema, nextConfig); + const currentModel = getValueByPath(fullConfig, "model.name"); + if (!isModelCompatibleWithDataset(currentModel, String(value))) { + nextConfig = setConfigValue(nextConfig, "model.name", "Auto"); + } + } return { ...exp, config_patch: nextConfig, @@ -300,6 +312,7 @@ function ExperimentCard({ {sectionFields.map((field) => ( ; disabled: boolean; field: AgentSchemaField; value: unknown; onChange: (value: unknown) => void; }) { + const compatibilityHint = fieldCompatibilityHint(field.path, config); + const boolDisabled = disabled || booleanDisabledForConfig(field.path, value, config); + const boolReason = booleanDisableReasonForConfig(field.path, value, config); + return (

{field.label}

@@ -350,16 +369,21 @@ function SchemaFieldControl({ {field.options.map((option) => { const meta = optionMeta(field.definition, option); + const optionDisabled = optionDisabledForConfig(field.definition, option, field.path, config); + const reason = optionDisableReasonForConfig(field.definition, option, field.path, config); return ( - - {optionLabel(field.definition, option)} - {typeof meta.badge === "string" && {meta.badge}} + + + {optionLabel(field.definition, option)} + {typeof meta.badge === "string" && {meta.badge}} + + {reason && {reason}} ); @@ -381,7 +405,7 @@ function SchemaFieldControl({ {field.type === "bool" && (
{formatFieldValue(value)} - +
)} {field.type !== "select" && field.type !== "number" && field.type !== "bool" && ( @@ -392,6 +416,8 @@ function SchemaFieldControl({ onChange={(event) => onChange(event.target.value)} /> )} + {compatibilityHint &&

{compatibilityHint}

} + {boolReason &&

{boolReason}

}
); } diff --git a/apps/frontend/src/features/agent/schema.ts b/apps/frontend/src/features/agent/schema.ts index aa558c3..230e20b 100644 --- a/apps/frontend/src/features/agent/schema.ts +++ b/apps/frontend/src/features/agent/schema.ts @@ -9,6 +9,14 @@ import { setValueByPath, type SchemaNode, } from "../simulation/utils"; +import { + dynamicBooleanState, + dynamicOptionState, + fieldHint, + optionLabel as compatibilityOptionLabel, + optionMeta as compatibilityOptionMeta, + staticOptionDisabled, +} from "../config/compatibility"; const META_KEYS = new Set(["role", "depends_on", "hidden", "ui"]); const AGENT_SECTIONS = new Set(["dataset", "model", "federated", "compression", "privacy"]); @@ -44,19 +52,49 @@ export function fieldLabel(key: string, definition: SchemaNode): string { } export function optionMeta(definition: SchemaNode, option: unknown): Record { - const options = fieldUi(definition).options; - if (!isRecord(options)) return {}; - const meta = options[String(option)]; - return isRecord(meta) ? meta : {}; + return compatibilityOptionMeta(definition, option); } export function optionLabel(definition: SchemaNode, option: unknown): string { - const meta = optionMeta(definition, option); - return typeof meta.label === "string" && meta.label.trim() ? meta.label : String(option); + return compatibilityOptionLabel(definition, option); } export function optionDisabled(definition: SchemaNode, option: unknown): boolean { - return optionMeta(definition, option).disabled === true; + return staticOptionDisabled(definition, option); +} + +export function optionDisabledForConfig( + definition: SchemaNode, + option: unknown, + path: string, + config: Record, +): boolean { + return optionDisabled(definition, option) || dynamicOptionState(path, option, config).disabled; +} + +export function optionDisableReasonForConfig( + definition: SchemaNode, + option: unknown, + path: string, + config: Record, +): string | null { + const meta = optionMeta(definition, option); + if (optionDisabled(definition, option)) { + return typeof meta.description === "string" ? meta.description : "Not executable yet."; + } + return dynamicOptionState(path, option, config).reason ?? null; +} + +export function booleanDisabledForConfig(path: string, value: unknown, config: Record): boolean { + return dynamicBooleanState(path, value, config).disabled; +} + +export function booleanDisableReasonForConfig(path: string, value: unknown, config: Record): string | null { + return dynamicBooleanState(path, value, config).reason ?? null; +} + +export function fieldCompatibilityHint(path: string, config: Record): string | null { + return fieldHint(path, config); } function isAgentVisible(definition: SchemaNode): boolean { @@ -175,8 +213,10 @@ export function validateDisabledOptions(config: Record, schema: for (const field of collectAgentSchemaFields(schema, { featuredOnly: false })) { if (field.type !== "select") continue; const value = getValueByPath(config, field.path); - if (value !== undefined && optionDisabled(field.definition, value)) { - errors.push(`${field.label}: ${optionLabel(field.definition, value)} is experimental and not executable yet.`); + if (value !== undefined && optionDisabledForConfig(field.definition, value, field.path, config)) { + const reason = optionDisableReasonForConfig(field.definition, value, field.path, config); + const suffix = reason ? ` ${reason}` : " It is not compatible with the current configuration."; + errors.push(`${field.label}: ${optionLabel(field.definition, value)} is unavailable.${suffix}`); } } return errors; diff --git a/apps/frontend/src/features/agent/useAgentController.ts b/apps/frontend/src/features/agent/useAgentController.ts index 8e0dc25..9ba4108 100644 --- a/apps/frontend/src/features/agent/useAgentController.ts +++ b/apps/frontend/src/features/agent/useAgentController.ts @@ -14,7 +14,7 @@ import type { AgentPageProps, AgentWorkflowStep, AgentPlanDraft } from "../../pa import { toErrorMessage } from "../simulation/utils"; import { removeValueByPath, setConfigValue } from "./schema"; -const DEFAULT_GOAL = "Compare CIFAR-10 non-IID with alpha=0.1, 0.3, 0.5"; +const DEFAULT_GOAL = "Compare non-IID alpha=0.1, 0.3, 0.5 on the selected dataset"; const POLL_INTERVAL_MS = 1500; function makeDefaultJobName(): string { @@ -277,7 +277,7 @@ export function useAgentController(): AgentPageProps { const presets = useMemo( () => [ - "Compare CIFAR-10 non-IID with alpha=0.1, 0.3, 0.5", + "Compare non-IID alpha=0.1, 0.3, 0.5 on the selected dataset", "Compare 10 clients vs 20 clients with FedAvg", "Test training rounds 10, 20, 50 on accuracy", ], diff --git a/apps/frontend/src/features/config/compatibility.ts b/apps/frontend/src/features/config/compatibility.ts new file mode 100644 index 0000000..0610709 --- /dev/null +++ b/apps/frontend/src/features/config/compatibility.ts @@ -0,0 +1,199 @@ +import { getValueByPath, isRecord, type SchemaNode } from "../simulation/utils"; + +const DATASET_MODEL_COMPATIBILITY: Record = { + MNIST: ["Mclr_Logistic", "LeNet", "DNN"], + EMNIST: ["Mclr_Logistic", "LeNet", "DNN"], + FEMNIST: ["Mclr_Logistic", "LeNet", "DNN"], + "Fashion-MNIST": ["Mclr_Logistic", "LeNet", "DNN"], + "CIFAR-10": ["Mclr_Logistic", "FedAvgCNN", "DNN", "ResNet18", "AlexNet", "MobileNet", "GoogleNet"], + "CIFAR-100": ["Mclr_Logistic", "FedAvgCNN", "DNN", "ResNet18", "AlexNet", "MobileNet", "GoogleNet"], + "AG News": ["TextLogistic", "TextDNN", "TextCNN"], + "Sogou News": ["TextLogistic", "TextDNN", "TextCNN"], + "Tiny-ImageNet": ["Mclr_Logistic", "FedAvgCNN", "DNN", "ResNet18", "AlexNet", "MobileNet", "GoogleNet"], + Country211: ["FedAvgCNN", "DNN", "ResNet18", "ResNet34", "AlexNet", "MobileNet", "GoogleNet"], + Flowers102: ["FedAvgCNN", "DNN", "ResNet18", "ResNet34", "AlexNet", "MobileNet", "GoogleNet"], + GTSRB: ["FedAvgCNN", "DNN", "ResNet18", "ResNet34", "AlexNet", "MobileNet", "GoogleNet"], + Shakespeare: ["CharLSTM"], + "Stanford Cars": ["FedAvgCNN", "DNN", "ResNet18", "ResNet34", "AlexNet", "MobileNet", "GoogleNet"], + COVIDx: ["FedAvgCNN", "DNN", "ResNet18", "ResNet34", "AlexNet", "MobileNet", "GoogleNet"], + Kvasir: ["FedAvgCNN", "DNN", "ResNet18", "ResNet34", "AlexNet", "MobileNet", "GoogleNet"], +}; + +const DATASET_DEFAULT_MODELS: Record = { + MNIST: "LeNet", + EMNIST: "LeNet", + FEMNIST: "LeNet", + "Fashion-MNIST": "LeNet", + "CIFAR-10": "FedAvgCNN", + "CIFAR-100": "FedAvgCNN", + "AG News": "TextDNN", + "Sogou News": "TextDNN", + "Tiny-ImageNet": "FedAvgCNN", + Country211: "MobileNet", + Flowers102: "MobileNet", + GTSRB: "MobileNet", + Shakespeare: "CharLSTM", + "Stanford Cars": "MobileNet", + COVIDx: "MobileNet", + Kvasir: "MobileNet", +}; + +const LINEAR_AGGREGATIONS = new Set(["fedavg", "weighted_avg", "simple_avg"]); +const EQUAL_WEIGHT_AGGREGATIONS = new Set(["fedavg", "simple_avg"]); + +export type OptionState = { + disabled: boolean; + reason?: string; +}; + +function boolAt(properties: Record, path: string): boolean { + return getValueByPath(properties, path) === true; +} + +function stringAt(properties: Record, path: string, fallback = ""): string { + const value = getValueByPath(properties, path); + return typeof value === "string" ? value : fallback; +} + +export function optionMeta(definition: SchemaNode, option: unknown): Record { + const ui = definition.ui; + if (!isRecord(ui)) return {}; + const options = ui.options; + if (!isRecord(options)) return {}; + const meta = options[String(option)]; + return isRecord(meta) ? meta : {}; +} + +export function optionLabel(definition: SchemaNode, option: unknown): string { + const meta = optionMeta(definition, option); + return typeof meta.label === "string" && meta.label.trim() ? meta.label : String(option); +} + +export function staticOptionDisabled(definition: SchemaNode, option: unknown): boolean { + return optionMeta(definition, option).disabled === true; +} + +export function compatibleModelsForDataset(datasetName: string): string[] { + return DATASET_MODEL_COMPATIBILITY[datasetName] ?? []; +} + +export function defaultModelForDataset(datasetName: string): string | null { + return DATASET_DEFAULT_MODELS[datasetName] ?? null; +} + +export function isModelCompatibleWithDataset(modelName: unknown, datasetName: string): boolean { + if (modelName === "Auto") return true; + if (typeof modelName !== "string") return false; + const models = compatibleModelsForDataset(datasetName); + return models.length === 0 || models.includes(modelName); +} + +export function dynamicOptionState( + path: string, + option: unknown, + properties: Record, +): OptionState { + const optionText = String(option); + const aggregation = stringAt(properties, "federated.aggregation", "fedavg"); + const ckksEnabled = boolAt(properties, "privacy.homomorphic_encryption.enable"); + const dpEnabled = boolAt(properties, "privacy.differential_privacy.enable"); + const compressionEnabled = boolAt(properties, "compression.sparsification.enable"); + const secureAggEnabled = boolAt(properties, "privacy.secure_aggregation.enable"); + + if (path === "model.name") { + const datasetName = stringAt(properties, "dataset.name", "CIFAR-10"); + if (!isModelCompatibleWithDataset(optionText, datasetName)) { + return { + disabled: true, + reason: `${optionText} is not compatible with ${datasetName}.`, + }; + } + } + + if (path === "federated.aggregation") { + if ((ckksEnabled || compressionEnabled) && !LINEAR_AGGREGATIONS.has(optionText)) { + return { + disabled: true, + reason: "CKKS and compression only support linear aggregation.", + }; + } + if (secureAggEnabled && !EQUAL_WEIGHT_AGGREGATIONS.has(optionText)) { + return { + disabled: true, + reason: "Secure aggregation masking requires equal-weight aggregation.", + }; + } + } + + if (path === "privacy.homomorphic_encryption.method" && optionText !== "ckks") { + return { disabled: true, reason: "Only CKKS is currently executable." }; + } + + if (path === "compression.sparsification.method") { + if (ckksEnabled) return { disabled: true, reason: "Turn off CKKS before enabling compression." }; + if (secureAggEnabled) return { disabled: true, reason: "Turn off secure aggregation before enabling compression." }; + } + + if (path === "privacy.differential_privacy.enable" && dpEnabled) { + return { disabled: false }; + } + + void aggregation; + return { disabled: false }; +} + +export function dynamicBooleanState(path: string, value: unknown, properties: Record): OptionState { + if (value === true) return { disabled: false }; + + const aggregation = stringAt(properties, "federated.aggregation", "fedavg"); + const ckksEnabled = boolAt(properties, "privacy.homomorphic_encryption.enable"); + const compressionEnabled = boolAt(properties, "compression.sparsification.enable"); + const secureAggEnabled = boolAt(properties, "privacy.secure_aggregation.enable"); + + if (path === "privacy.homomorphic_encryption.enable") { + if (!LINEAR_AGGREGATIONS.has(aggregation)) { + return { disabled: true, reason: "CKKS requires fedavg, weighted_avg, or simple_avg." }; + } + if (compressionEnabled) return { disabled: true, reason: "Turn off compression before enabling CKKS." }; + if (secureAggEnabled) return { disabled: true, reason: "Turn off secure aggregation before enabling CKKS." }; + } + + if (path === "compression.sparsification.enable") { + if (!LINEAR_AGGREGATIONS.has(aggregation)) { + return { disabled: true, reason: "Compression requires fedavg, weighted_avg, or simple_avg." }; + } + if (ckksEnabled) return { disabled: true, reason: "Turn off CKKS before enabling compression." }; + if (secureAggEnabled) return { disabled: true, reason: "Turn off secure aggregation before enabling compression." }; + } + + if (path === "privacy.secure_aggregation.enable") { + if (!EQUAL_WEIGHT_AGGREGATIONS.has(aggregation)) { + return { disabled: true, reason: "Secure aggregation requires fedavg or simple_avg." }; + } + if (ckksEnabled) return { disabled: true, reason: "Turn off CKKS before enabling secure aggregation." }; + if (compressionEnabled) return { disabled: true, reason: "Turn off compression before enabling secure aggregation." }; + } + + return { disabled: false }; +} + +export function fieldHint(path: string, properties: Record): string | null { + if (path === "model.name") { + const datasetName = stringAt(properties, "dataset.name", "CIFAR-10"); + const models = compatibleModelsForDataset(datasetName); + if (models.length === 0) return null; + const defaultModel = defaultModelForDataset(datasetName); + const defaultText = defaultModel ? ` Default: ${defaultModel}.` : ""; + return `${datasetName} supports: Auto, ${models.join(", ")}.${defaultText}`; + } + + if (path === "federated.aggregation") { + const ckksEnabled = boolAt(properties, "privacy.homomorphic_encryption.enable"); + const compressionEnabled = boolAt(properties, "compression.sparsification.enable"); + const secureAggEnabled = boolAt(properties, "privacy.secure_aggregation.enable"); + if (secureAggEnabled) return "Secure aggregation masking allows fedavg or simple_avg."; + if (ckksEnabled || compressionEnabled) return "CKKS and compression allow fedavg, weighted_avg, or simple_avg."; + } + + return null; +} diff --git a/apps/frontend/src/features/config/render-schema-section.tsx b/apps/frontend/src/features/config/render-schema-section.tsx index f12f556..1097e65 100644 --- a/apps/frontend/src/features/config/render-schema-section.tsx +++ b/apps/frontend/src/features/config/render-schema-section.tsx @@ -6,6 +6,15 @@ import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from ". import { Switch } from "../../components/ui/switch"; import type { NodeRole } from "../../pages/types"; import { getValueByPath, isFieldDefinition, isRecord, normalizeListInt, type SchemaNode } from "../simulation/utils"; +import { + dynamicBooleanState, + dynamicOptionState, + fieldHint, + isModelCompatibleWithDataset, + optionLabel, + optionMeta, + staticOptionDisabled, +} from "./compatibility"; function checkDependency(properties: Record, dependsOn?: string): boolean { if (!dependsOn) return true; @@ -70,17 +79,22 @@ export function renderSchemaSection( const value = getValueByPath(properties, fullPath); if (definition.type === "bool") { const label = labelize(key); + const state = dynamicBooleanState(fullPath, value, properties); blocks.push(
- {label} - onChangeProperty(fullPath, checked)} - aria-label={label} - /> +
+ {label} + onChangeProperty(fullPath, checked)} + aria-label={label} + /> +
+ {state.reason &&

{state.reason}

}
, ); continue; @@ -88,6 +102,7 @@ export function renderSchemaSection( if (definition.type === "select") { const options = Array.isArray(definition.options) ? definition.options : []; + const hint = fieldHint(fullPath, properties); blocks.push(

{labelize(key)}

@@ -95,20 +110,50 @@ export function renderSchemaSection( value={String(value ?? "")} onValueChange={(nextValue) => { const matched = options.find((option) => String(option) === nextValue); - onChangeProperty(fullPath, matched ?? nextValue); + const next = matched ?? nextValue; + onChangeProperty(fullPath, next); + if (fullPath === "dataset.name") { + const currentModel = getValueByPath(properties, "model.name"); + if (!isModelCompatibleWithDataset(currentModel, String(next))) { + onChangeProperty("model.name", "Auto"); + } + } }} > - {options.map((option) => ( - - {String(option)} - - ))} + {options.map((option) => { + const staticDisabled = staticOptionDisabled(definition, option); + const dynamicState = dynamicOptionState(fullPath, option, properties); + const disabled = staticDisabled || dynamicState.disabled; + const meta = optionMeta(definition, option); + const reason = staticDisabled + ? typeof meta.description === "string" + ? meta.description + : "Not executable yet." + : dynamicState.reason; + return ( + + + + {optionLabel(definition, option)} + {typeof meta.badge === "string" && {meta.badge}} + + {reason && {reason}} + + + ); + })} + {hint &&

{hint}

}
, ); continue; diff --git a/config_schema.yaml b/config_schema.yaml index 4beefb1..2a54f83 100644 --- a/config_schema.yaml +++ b/config_schema.yaml @@ -5,7 +5,7 @@ dataset: group: "Data" name: type: "select" - options: ["CIFAR-10", "CIFAR-100", "MNIST", "Fashion-MNIST"] + options: ["MNIST", "EMNIST", "FEMNIST", "Fashion-MNIST", "CIFAR-10", "CIFAR-100", "AG News", "Sogou News", "Tiny-ImageNet", "Country211", "Flowers102", "GTSRB", "Shakespeare", "Stanford Cars", "COVIDx", "Kvasir"] default: "CIFAR-10" ui: label: "Dataset" @@ -55,8 +55,8 @@ model: group: "Model" name: type: "select" - options: ["CNN", "LeNet", "ResNet"] - default: "CNN" + options: ["Auto", "Mclr_Logistic", "LeNet", "DNN", "FedAvgCNN", "ResNet18", "ResNet34", "ResNet50", "AlexNet", "MobileNet", "GoogleNet", "TextLogistic", "TextDNN", "TextCNN", "CharLSTM"] + default: "Auto" ui: label: "Model" featured: true @@ -64,7 +64,7 @@ model: prompt_hint: "Neural network architecture for every client." input_shape: type: "list_int" - default: [32, 32, 3] + default: [3, 32, 32] ui: label: "Input shape" num_classes: @@ -125,7 +125,7 @@ federated: prompt_hint: "Optimizer learning rate for local training." aggregation: type: "select" - options: ["fedavg", "fedprox", "scaffold"] + options: ["fedavg", "weighted_avg", "simple_avg", "fedprox", "scaffold"] default: "fedavg" ui: label: "Aggregation" @@ -136,6 +136,12 @@ federated: fedavg: label: "FedAvg" description: "Executable default aggregation strategy." + weighted_avg: + label: "Weighted Avg" + description: "Executable sample-weighted averaging strategy." + simple_avg: + label: "Simple Avg" + description: "Executable unweighted averaging strategy." fedprox: label: "FedProx" badge: "experimental" @@ -164,8 +170,32 @@ compression: type: "bool" default: false ui: - label: "Enable sparsification" + label: "Enable compression" group: "Advanced" + method: + type: "select" + options: ["global_topk", "random_k", "threshold", "sign", "quant_int8"] + default: "global_topk" + depends_on: "compression.sparsification.enable == true" + ui: + label: "Compression method" + group: "Advanced" + options: + global_topk: + label: "Global Top-K" + description: "Keep the largest update values globally." + random_k: + label: "Random-K" + description: "Keep a random subset of update values." + threshold: + label: "Threshold" + description: "Drop update values below an absolute threshold." + sign: + label: "Sign" + description: "Send sign-only updates scaled by mean magnitude." + quant_int8: + label: "Int8 Quantization" + description: "Quantize updates to int8 and dequantize before aggregation." ratio: type: "number" default: 0.5 @@ -173,7 +203,16 @@ compression: max: 1.0 depends_on: "compression.sparsification.enable == true" ui: - label: "Top-K ratio" + label: "Keep ratio" + group: "Advanced" + threshold: + type: "number" + default: 0.001 + min: 0 + step: 0.0001 + depends_on: "compression.sparsification.enable == true" + ui: + label: "Threshold" group: "Advanced" privacy: @@ -188,6 +227,14 @@ privacy: ui: label: "Enable CKKS" group: "Advanced" + method: + type: "select" + options: ["ckks"] + default: "ckks" + depends_on: "privacy.homomorphic_encryption.enable == true" + ui: + label: "Encryption method" + group: "Advanced" poly_modulus_degree: type: "select" options: [4096, 8192, 16384] @@ -196,6 +243,47 @@ privacy: ui: label: "Polynomial modulus degree" group: "Advanced" + differential_privacy: + enable: + type: "bool" + default: false + ui: + label: "Enable DP" + group: "Advanced" + clipping_norm: + type: "number" + default: 1.0 + min: 0.0001 + step: 0.1 + depends_on: "privacy.differential_privacy.enable == true" + ui: + label: "Clipping norm" + group: "Advanced" + noise_multiplier: + type: "number" + default: 0.0 + min: 0 + step: 0.1 + depends_on: "privacy.differential_privacy.enable == true" + ui: + label: "Noise multiplier" + group: "Advanced" + secure_aggregation: + enable: + type: "bool" + default: false + ui: + label: "Enable secure aggregation" + group: "Advanced" + mask_std: + type: "number" + default: 1.0 + min: 0 + step: 0.1 + depends_on: "privacy.secure_aggregation.enable == true" + ui: + label: "Mask std" + group: "Advanced" system: hidden: true diff --git a/configs/simulation_config.yaml b/configs/simulation_config.yaml index 073f5b0..fc67d9c 100644 --- a/configs/simulation_config.yaml +++ b/configs/simulation_config.yaml @@ -5,11 +5,11 @@ dataset: alpha: 0.5 num_clients: 3 model: - name: CNN + name: Auto input_shape: + - 3 - 32 - 32 - - 3 num_classes: 10 federated: num_rounds: 10 @@ -22,11 +22,21 @@ federated: compression: sparsification: enable: false + method: global_topk ratio: 0.5 + threshold: 0.001 privacy: homomorphic_encryption: enable: false + method: ckks poly_modulus_degree: 8192 + differential_privacy: + enable: false + clipping_norm: 1.0 + noise_multiplier: 0.0 + secure_aggregation: + enable: false + mask_std: 1.0 system: mode: simulation node_role: server diff --git a/libs/fl_core/compression/__init__.py b/libs/fl_core/compression/__init__.py index e69de29..9311568 100644 --- a/libs/fl_core/compression/__init__.py +++ b/libs/fl_core/compression/__init__.py @@ -0,0 +1,19 @@ +from .sparsification import ( + CompressionFactory, + CompressionStrategy, + GlobalTopKSparsifier, + QuantInt8Compressor, + RandomKSparsifier, + SignCompressor, + ThresholdSparsifier, +) + +__all__ = [ + "CompressionFactory", + "CompressionStrategy", + "GlobalTopKSparsifier", + "RandomKSparsifier", + "ThresholdSparsifier", + "SignCompressor", + "QuantInt8Compressor", +] diff --git a/libs/fl_core/compression/sparsification.py b/libs/fl_core/compression/sparsification.py index 453d84a..d8ef63b 100644 --- a/libs/fl_core/compression/sparsification.py +++ b/libs/fl_core/compression/sparsification.py @@ -1,46 +1,167 @@ -import torch -from typing import Dict +from __future__ import annotations + import logging +from abc import ABC, abstractmethod +from typing import Any, Dict + +import torch + + +def _is_trainable_float(name: str, tensor: torch.Tensor) -> bool: + return ( + tensor.is_floating_point() + and tensor.dim() > 0 + and "num_batches_tracked" not in name + and "running_" not in name + ) -class GlobalTopKSparsifier: - def __init__(self, ratio: float = 0.5): - self.ratio = ratio - self.logger = logging.getLogger("Sparsifier") +class CompressionStrategy(ABC): + method = "none" + + @abstractmethod def sparsify(self, update_dict: Dict[str, torch.Tensor]) -> Dict[str, torch.Tensor]: + pass + + +class GlobalTopKSparsifier(CompressionStrategy): + method = "global_topk" + def __init__(self, ratio: float = 0.5, **_: Any): + self.ratio = ratio + self.logger = logging.getLogger("GlobalTopKSparsifier") + + def sparsify(self, update_dict: Dict[str, torch.Tensor]) -> Dict[str, torch.Tensor]: if self.ratio >= 1.0: return update_dict - # 1. 收集所有参与训练的参数(排除BN层统计量和标量) - trainable_tensors = [] - for name, tensor in update_dict.items(): - if tensor.dim() > 0 and 'num_batches_tracked' not in name and 'running_' not in name: - trainable_tensors.append(tensor.abs().view(-1)) - + trainable_tensors = [ + tensor.abs().view(-1) + for name, tensor in update_dict.items() + if _is_trainable_float(name, tensor) + ] if not trainable_tensors: return update_dict - # 2. 全局拼接并计算阈值 all_abs_params = torch.cat(trainable_tensors) k = int(all_abs_params.numel() * self.ratio) - if k == 0: - return {k: torch.zeros_like(v) for k, v in update_dict.items()} + return { + name: torch.zeros_like(tensor) if _is_trainable_float(name, tensor) else tensor + for name, tensor in update_dict.items() + } - # 获取第 k 大的值作为阈值 threshold = torch.kthvalue(all_abs_params, all_abs_params.numel() - k + 1).values.item() + sparse_update = {} + for name, tensor in update_dict.items(): + if not _is_trainable_float(name, tensor): + sparse_update[name] = tensor + continue + sparse_update[name] = tensor * (torch.abs(tensor) >= threshold).to(tensor.dtype) + return sparse_update + + +class RandomKSparsifier(CompressionStrategy): + method = "random_k" + + def __init__(self, ratio: float = 0.5, **_: Any): + self.ratio = ratio + self.logger = logging.getLogger("RandomKSparsifier") + + def sparsify(self, update_dict: Dict[str, torch.Tensor]) -> Dict[str, torch.Tensor]: + if self.ratio >= 1.0: + return update_dict + sparse_update = {} + for name, tensor in update_dict.items(): + if not _is_trainable_float(name, tensor): + sparse_update[name] = tensor + continue + flat = tensor.view(-1) + k = int(flat.numel() * self.ratio) + if k <= 0: + sparse_update[name] = torch.zeros_like(tensor) + continue + mask = torch.zeros(flat.numel(), device=tensor.device, dtype=torch.bool) + indices = torch.randperm(flat.numel(), device=tensor.device)[:k] + mask[indices] = True + sparse_update[name] = (flat * mask.to(flat.dtype)).view_as(tensor) + return sparse_update + + +class ThresholdSparsifier(CompressionStrategy): + method = "threshold" - # 3. 应用掩码 + def __init__(self, threshold: float = 1e-3, **_: Any): + self.threshold = float(threshold) + self.ratio = 1.0 + self.logger = logging.getLogger("ThresholdSparsifier") + + def sparsify(self, update_dict: Dict[str, torch.Tensor]) -> Dict[str, torch.Tensor]: sparse_update = {} for name, tensor in update_dict.items(): - # 不处理统计量 - if tensor.dim() == 0 or 'num_batches_tracked' in name or 'running_' in name: + if not _is_trainable_float(name, tensor): sparse_update[name] = tensor continue + sparse_update[name] = tensor * (torch.abs(tensor) >= self.threshold).to(tensor.dtype) + return sparse_update + + +class SignCompressor(CompressionStrategy): + method = "sign" + + def __init__(self, **_: Any): + self.ratio = 1.0 + self.logger = logging.getLogger("SignCompressor") + + def sparsify(self, update_dict: Dict[str, torch.Tensor]) -> Dict[str, torch.Tensor]: + compressed = {} + for name, tensor in update_dict.items(): + if not _is_trainable_float(name, tensor): + compressed[name] = tensor + continue + mean_abs = tensor.abs().mean().clamp_min(1e-12) + compressed[name] = tensor.sign() * mean_abs + return compressed + + +class QuantInt8Compressor(CompressionStrategy): + method = "quant_int8" + + def __init__(self, **_: Any): + self.ratio = 1.0 + self.logger = logging.getLogger("QuantInt8Compressor") + + def sparsify(self, update_dict: Dict[str, torch.Tensor]) -> Dict[str, torch.Tensor]: + compressed = {} + for name, tensor in update_dict.items(): + if not _is_trainable_float(name, tensor): + compressed[name] = tensor + continue + scale = tensor.abs().max().clamp_min(1e-12) / 127.0 + quantized = torch.clamp(torch.round(tensor / scale), -127, 127) + compressed[name] = quantized * scale + return compressed + + +class CompressionFactory: + _strategies = { + "global_topk": GlobalTopKSparsifier, + "topk": GlobalTopKSparsifier, + "random_k": RandomKSparsifier, + "randomk": RandomKSparsifier, + "threshold": ThresholdSparsifier, + "sign": SignCompressor, + "quant_int8": QuantInt8Compressor, + "int8": QuantInt8Compressor, + } + + @classmethod + def create(cls, method: str = "global_topk", **kwargs: Any) -> CompressionStrategy: + key = method.lower() + if key not in cls._strategies: + raise ValueError(f"Unsupported compression method: {method}. Supported: {list(cls._strategies.keys())}") + return cls._strategies[key](**kwargs) - # 生成掩码 - mask = torch.abs(tensor) >= threshold - sparse_update[name] = tensor * mask.float() - - return sparse_update \ No newline at end of file + @classmethod + def supported_methods(cls) -> list[str]: + return ["global_topk", "random_k", "threshold", "sign", "quant_int8"] diff --git a/libs/fl_core/data/data_loader.py b/libs/fl_core/data/data_loader.py index 6b604bc..f6415ff 100644 --- a/libs/fl_core/data/data_loader.py +++ b/libs/fl_core/data/data_loader.py @@ -1,154 +1,404 @@ +from __future__ import annotations + +import csv +import hashlib +import json import os +from pathlib import Path +from typing import Any, Dict, Iterable, Tuple + +import numpy as np import torch import torchvision import torchvision.transforms as transforms -from torch.utils.data import DataLoader -import numpy as np -from typing import Tuple, Dict, Any +from PIL import Image +from torch.utils.data import DataLoader, Dataset + +from fl_core.simulation_registry import ( + DatasetSpec, + canonical_dataset_name, + get_dataset_spec, + list_dataset_names, +) + + +def _safe_dataset_class(name: str): + return getattr(torchvision.datasets, name, None) + + +def _image_transform(input_shape: tuple[int, int, int], mean: tuple[float, ...] | None = None, std: tuple[float, ...] | None = None): + channels, height, width = input_shape + steps: list[Any] = [] + if (height, width) not in {(28, 28), (32, 32)}: + steps.append(transforms.Resize((height, width))) + steps.append(transforms.Grayscale(num_output_channels=channels) if channels == 1 else transforms.Lambda(lambda image: image.convert("RGB"))) + steps.append(transforms.ToTensor()) + if mean is not None and std is not None: + steps.append(transforms.Normalize(mean, std)) + return transforms.Compose(steps) + + +def _hash_text_to_vector(text: str, feature_dim: int) -> np.ndarray: + vector = np.zeros(feature_dim, dtype=np.float32) + for token in text.lower().split(): + if not token: + continue + digest = hashlib.md5(token.encode("utf-8")).hexdigest() + vector[int(digest[:8], 16) % feature_dim] += 1.0 + norm = np.linalg.norm(vector) + if norm > 0: + vector /= norm + return vector + + +class CsvTextClassificationDataset(Dataset): + def __init__(self, path: Path, *, feature_dim: int): + if not path.exists(): + raise FileNotFoundError(f"Text dataset file not found: {path}") + self.samples: list[tuple[np.ndarray, int]] = [] + with path.open("r", encoding="utf-8", errors="replace", newline="") as fp: + reader = csv.reader(fp) + for row in reader: + if len(row) < 2: + continue + try: + label = int(row[0]) + except ValueError: + continue + if label > 0: + label -= 1 + text = " ".join(row[1:]) + self.samples.append((_hash_text_to_vector(text, feature_dim), label)) + if not self.samples: + raise ValueError(f"No text samples found in {path}") + + def __len__(self) -> int: + return len(self.samples) + + def __getitem__(self, index: int): + data, label = self.samples[index] + return torch.from_numpy(data), int(label) + + +class ShakespeareDataset(Dataset): + def __init__(self, path: Path, *, sequence_length: int = 80, vocab_size: int = 128): + if not path.exists(): + raise FileNotFoundError(f"Shakespeare text file not found: {path}") + text = path.read_text(encoding="utf-8", errors="replace") + encoded = np.array([min(ord(ch), vocab_size - 1) for ch in text], dtype=np.int64) + if len(encoded) <= sequence_length: + raise ValueError(f"Shakespeare file {path} is too short for sequence_length={sequence_length}") + self.x = [] + self.y = [] + for start in range(0, len(encoded) - sequence_length): + end = start + sequence_length + self.x.append(encoded[start:end]) + self.y.append(int(encoded[end])) + + def __len__(self) -> int: + return len(self.y) + + def __getitem__(self, index: int): + return torch.from_numpy(self.x[index]), self.y[index] + + +class LeafFEMNISTDataset(Dataset): + def __init__(self, root: Path, split: str): + files = sorted((root / split).glob("*.json")) + if not files: + raise FileNotFoundError( + f"FEMNIST LEAF files not found. Expected JSON files under {root / split}" + ) + data: list[np.ndarray] = [] + labels: list[int] = [] + for path in files: + loaded = json.loads(path.read_text(encoding="utf-8")) + user_data = loaded.get("user_data", {}) + for user_payload in user_data.values(): + xs = user_payload.get("x", []) + ys = user_payload.get("y", []) + for x, y in zip(xs, ys): + arr = np.asarray(x, dtype=np.float32).reshape(1, 28, 28) + data.append(arr) + labels.append(int(y)) + if not data: + raise ValueError(f"No FEMNIST samples found under {root / split}") + self.data = np.stack(data, axis=0) + self.labels = np.asarray(labels, dtype=np.int64) + + def __len__(self) -> int: + return len(self.labels) + + def __getitem__(self, index: int): + return torch.from_numpy(self.data[index]), int(self.labels[index]) + + +class TinyImageNetValDataset(Dataset): + def __init__(self, root: Path, transform, class_to_idx: dict[str, int] | None = None): + val_dir = root / "val" + annotation_path = val_dir / "val_annotations.txt" + if not annotation_path.exists(): + raise FileNotFoundError(f"Tiny-ImageNet val annotations not found: {annotation_path}") + if class_to_idx is None: + wnids = _read_tiny_imagenet_wnids(root) + class_to_idx = {wnid: idx for idx, wnid in enumerate(wnids)} + self.class_to_idx = class_to_idx + self.samples: list[tuple[Path, int]] = [] + with annotation_path.open("r", encoding="utf-8") as fp: + for line in fp: + parts = line.strip().split("\t") + if len(parts) < 2: + continue + image_name, wnid = parts[:2] + if wnid not in self.class_to_idx: + continue + self.samples.append((val_dir / "images" / image_name, self.class_to_idx[wnid])) + self.transform = transform + if not self.samples: + raise ValueError(f"No Tiny-ImageNet validation samples found under {val_dir}") + + def __len__(self) -> int: + return len(self.samples) + + def __getitem__(self, index: int): + path, label = self.samples[index] + image = Image.open(path) + if self.transform: + image = self.transform(image) + return image, label + + +def _read_tiny_imagenet_wnids(root: Path) -> list[str]: + wnids_path = root / "wnids.txt" + if wnids_path.exists(): + return [line.strip() for line in wnids_path.read_text(encoding="utf-8").splitlines() if line.strip()] + train_dir = root / "train" + if train_dir.exists(): + return sorted(path.name for path in train_dir.iterdir() if path.is_dir()) + raise FileNotFoundError(f"Tiny-ImageNet class list not found under {root}") class DatasetLoader: - def __init__(self, data_dir: str = "./datasets"): self.data_dir = data_dir - self.supported_datasets = ["CIFAR-10", "CIFAR-100", "MNIST", "Fashion-MNIST"] - + self.supported_datasets = list_dataset_names() os.makedirs(self.data_dir, exist_ok=True) - - self.dataset_configs = { + self.dataset_configs = self._build_dataset_configs() + + def _build_dataset_configs(self) -> dict[str, dict[str, Any]]: + return { + "MNIST": { + "dataset_class": torchvision.datasets.MNIST, + "style": "train_bool", + "transform": _image_transform((1, 28, 28), (0.1307,), (0.3081,)), + }, + "EMNIST": { + "dataset_class": torchvision.datasets.EMNIST, + "style": "emnist", + "split": "balanced", + "transform": _image_transform((1, 28, 28), (0.1751,), (0.3332,)), + }, + "Fashion-MNIST": { + "dataset_class": torchvision.datasets.FashionMNIST, + "style": "train_bool", + "transform": _image_transform((1, 28, 28), (0.2860,), (0.3530,)), + }, "CIFAR-10": { "dataset_class": torchvision.datasets.CIFAR10, - "num_classes": 10, - "input_shape": (3, 32, 32), - "transform": transforms.Compose([ - transforms.ToTensor(), - transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2023, 0.1994, 0.2010)) - ]) + "style": "train_bool", + "transform": _image_transform((3, 32, 32), (0.4914, 0.4822, 0.4465), (0.2023, 0.1994, 0.2010)), }, "CIFAR-100": { "dataset_class": torchvision.datasets.CIFAR100, - "num_classes": 100, - "input_shape": (3, 32, 32), - "transform": transforms.Compose([ - transforms.ToTensor(), - transforms.Normalize((0.5071, 0.4867, 0.4408), (0.2675, 0.2565, 0.2761)) - ]) + "style": "train_bool", + "transform": _image_transform((3, 32, 32), (0.5071, 0.4867, 0.4408), (0.2675, 0.2565, 0.2761)), }, - "MNIST": { - "dataset_class": torchvision.datasets.MNIST, - "num_classes": 10, - "input_shape": (1, 28, 28), - "transform": transforms.Compose([ - transforms.ToTensor(), - transforms.Normalize((0.1307,), (0.3081,)) - ]) + "Country211": { + "dataset_class": _safe_dataset_class("Country211"), + "style": "split", + "train_split": "train", + "test_split": "test", + "transform": _image_transform((3, 64, 64)), + }, + "Flowers102": { + "dataset_class": _safe_dataset_class("Flowers102"), + "style": "split", + "train_split": "train", + "test_split": "test", + "transform": _image_transform((3, 64, 64)), + }, + "GTSRB": { + "dataset_class": _safe_dataset_class("GTSRB"), + "style": "split", + "train_split": "train", + "test_split": "test", + "transform": _image_transform((3, 64, 64)), + }, + "Stanford Cars": { + "dataset_class": _safe_dataset_class("StanfordCars"), + "style": "split", + "train_split": "train", + "test_split": "test", + "transform": _image_transform((3, 64, 64)), }, - "Fashion-MNIST": { - "dataset_class": torchvision.datasets.FashionMNIST, - "num_classes": 10, - "input_shape": (1, 28, 28), - "transform": transforms.Compose([ - transforms.ToTensor(), - transforms.Normalize((0.2860,), (0.3530,)) - ]) - } } - - def load_dataset(self, dataset_name: str) -> Tuple[np.ndarray, np.ndarray, np.ndarray, np.ndarray]: - if dataset_name not in self.supported_datasets: - raise ValueError(f"不支持的数据集: {dataset_name}. 支持的数据集: {self.supported_datasets}") + def load_dataset(self, dataset_name: str) -> Tuple[np.ndarray, np.ndarray, np.ndarray, np.ndarray]: + dataset_name = canonical_dataset_name(dataset_name) + spec = get_dataset_spec(dataset_name) + if spec is None: + raise ValueError(f"Unsupported dataset: {dataset_name}. Supported datasets: {self.supported_datasets}") already_cached = self.check_dataset_exists(dataset_name) if already_cached: print(f"DATASET_READY: {dataset_name} (cached at {self.data_dir})", flush=True) else: - print( - f"DATASET_DOWNLOAD_START: {dataset_name} (first-time download, this may take several minutes)", - flush=True, - ) + print(f"DATASET_DOWNLOAD_START: {dataset_name} (first-time load, this may take several minutes)", flush=True) - config = self.dataset_configs[dataset_name] - dataset_class = config["dataset_class"] - transform = config["transform"] + train_dataset, test_dataset = self._create_datasets(spec) - try: - # 加载训练集 - train_dataset = dataset_class( - root=self.data_dir, - train=True, - download=True, - transform=transform + if not already_cached: + print(f"DATASET_DOWNLOAD_DONE: {dataset_name}", flush=True) + print(f"DATASET_LOADED: {dataset_name} train={len(train_dataset)} test={len(test_dataset)}", flush=True) + + train_data, train_labels = self._dataset_to_numpy(train_dataset) + test_data, test_labels = self._dataset_to_numpy(test_dataset) + return train_data, train_labels, test_data, test_labels + + def _create_datasets(self, spec: DatasetSpec): + if spec.loader == "torchvision": + return self._create_torchvision_datasets(spec) + if spec.loader == "tiny_imagenet": + return self._create_tiny_imagenet_datasets(spec) + if spec.loader == "local_image_folder": + return self._create_local_image_folder_datasets(spec) + if spec.loader == "leaf_femnist": + root = Path(self.data_dir) / spec.name + return LeafFEMNISTDataset(root, "train"), LeafFEMNISTDataset(root, "test") + if spec.loader == "csv_text": + root = Path(self.data_dir) / spec.name.replace(" ", "_") + if not root.exists(): + root = Path(self.data_dir) / spec.name + feature_dim = int(spec.metadata.get("feature_dim", spec.input_shape[0])) + return ( + CsvTextClassificationDataset(root / "train.csv", feature_dim=feature_dim), + CsvTextClassificationDataset(root / "test.csv", feature_dim=feature_dim), ) + if spec.loader == "shakespeare": + root = Path(self.data_dir) / spec.name + seq_len = int(spec.metadata.get("sequence_length", spec.input_shape[0])) + vocab_size = int(spec.metadata.get("vocab_size", spec.num_classes)) + return ( + ShakespeareDataset(root / "train.txt", sequence_length=seq_len, vocab_size=vocab_size), + ShakespeareDataset(root / "test.txt", sequence_length=seq_len, vocab_size=vocab_size), + ) + raise ValueError(f"Unsupported dataset loader: {spec.loader}") - # 加载测试集 - test_dataset = dataset_class( - root=self.data_dir, - train=False, - download=True, - transform=transform + def _create_torchvision_datasets(self, spec: DatasetSpec): + config = self.dataset_configs.get(spec.name) + if not config: + raise ValueError(f"Missing torchvision config for dataset: {spec.name}") + dataset_class = config.get("dataset_class") + if dataset_class is None: + raise ValueError(f"torchvision does not provide dataset class for {spec.name}") + transform = config["transform"] + style = config["style"] + if style == "train_bool": + return ( + dataset_class(root=self.data_dir, train=True, download=True, transform=transform), + dataset_class(root=self.data_dir, train=False, download=True, transform=transform), ) + if style == "emnist": + split = config.get("split", "balanced") + return ( + dataset_class(root=self.data_dir, split=split, train=True, download=True, transform=transform), + dataset_class(root=self.data_dir, split=split, train=False, download=True, transform=transform), + ) + if style == "split": + return ( + dataset_class(root=self.data_dir, split=config["train_split"], download=True, transform=transform), + dataset_class(root=self.data_dir, split=config["test_split"], download=True, transform=transform), + ) + raise ValueError(f"Unsupported torchvision loader style: {style}") - if not already_cached: - print(f"DATASET_DOWNLOAD_DONE: {dataset_name}", flush=True) - - print( - f"DATASET_LOADED: {dataset_name} train={len(train_dataset)} test={len(test_dataset)}", - flush=True, + def _create_tiny_imagenet_datasets(self, spec: DatasetSpec): + root = Path(self.data_dir) / "Tiny-ImageNet" + if not root.exists(): + root = Path(self.data_dir) / "tiny-imagenet-200" + transform = _image_transform(spec.input_shape) + train_dir = root / "train" + if not train_dir.exists(): + raise FileNotFoundError(f"Tiny-ImageNet train directory not found: {train_dir}") + train_dataset = torchvision.datasets.ImageFolder(str(train_dir), transform=transform) + test_dataset = TinyImageNetValDataset(root, transform, train_dataset.class_to_idx) + return train_dataset, test_dataset + + def _create_local_image_folder_datasets(self, spec: DatasetSpec): + root = Path(self.data_dir) / spec.name + if not root.exists(): + root = Path(self.data_dir) / spec.name.lower() + train_dir = root / "train" + test_dir = root / "test" + if not train_dir.exists() or not test_dir.exists(): + raise FileNotFoundError( + f"{spec.name} expects ImageFolder layout under {root}: train//*.jpg and test//*.jpg" ) - - # 转换为numpy数组 - train_data, train_labels = self._dataset_to_numpy(train_dataset) - test_data, test_labels = self._dataset_to_numpy(test_dataset) - - return train_data, train_labels, test_data, test_labels - - except Exception as e: - print(f"加载数据集 {dataset_name} 时出错: {str(e)}") - raise - + transform = _image_transform(spec.input_shape) + return ( + torchvision.datasets.ImageFolder(str(train_dir), transform=transform), + torchvision.datasets.ImageFolder(str(test_dir), transform=transform), + ) + def _dataset_to_numpy(self, dataset) -> Tuple[np.ndarray, np.ndarray]: data_list = [] labels_list = [] - dataloader = DataLoader(dataset, batch_size=1000, shuffle=False) - for batch_data, batch_labels in dataloader: data_list.append(batch_data.numpy()) labels_list.append(batch_labels.numpy()) - data = np.concatenate(data_list, axis=0) labels = np.concatenate(labels_list, axis=0) - return data, labels - + def get_dataset_info(self, dataset_name: str) -> Dict[str, Any]: - if dataset_name not in self.supported_datasets: - raise ValueError(f"不支持的数据集: {dataset_name}") - - config = self.dataset_configs[dataset_name] - return { - "name": dataset_name, - "num_classes": config["num_classes"], - "input_shape": config["input_shape"], - "supported": True + dataset_name = canonical_dataset_name(dataset_name) + spec = get_dataset_spec(dataset_name) + if spec is None: + raise ValueError(f"Unsupported dataset: {dataset_name}") + info = { + "name": spec.name, + "num_classes": spec.num_classes, + "input_shape": spec.input_shape, + "modality": spec.modality, + "task_type": spec.task_type, + "default_model": spec.default_model, + "compatible_models": list(spec.compatible_models), + "supported": True, } - + info.update(spec.metadata) + return info + def list_supported_datasets(self) -> list: return self.supported_datasets.copy() - + def check_dataset_exists(self, dataset_name: str) -> bool: + dataset_name = canonical_dataset_name(dataset_name) if dataset_name not in self.supported_datasets: return False - + dataset_dir_map = { "CIFAR-10": "cifar-10-batches-py", - "CIFAR-100": "cifar-100-python", + "CIFAR-100": "cifar-100-python", "MNIST": "MNIST", - "Fashion-MNIST": "FashionMNIST" + "EMNIST": "EMNIST", + "Fashion-MNIST": "FashionMNIST", + "Tiny-ImageNet": "Tiny-ImageNet", + "FEMNIST": "FEMNIST", + "AG News": "AG_News", + "Sogou News": "Sogou_News", + "Shakespeare": "Shakespeare", } - - dataset_dir = os.path.join(self.data_dir, dataset_dir_map[dataset_name]) + dataset_dir = os.path.join(self.data_dir, dataset_dir_map.get(dataset_name, dataset_name)) return os.path.exists(dataset_dir) @@ -157,15 +407,13 @@ def __init__(self, config: Dict[str, Any]): self.config = config self.data_dir = config.get("data_dir", "./datasets") self.loader = DatasetLoader(self.data_dir) - + def load_dataset(self, dataset_name: str = None) -> Tuple[np.ndarray, np.ndarray, np.ndarray, np.ndarray]: if dataset_name is None: dataset_name = self.config.get("name", "CIFAR-10") - return self.loader.load_dataset(dataset_name) - + def get_dataset_info(self, dataset_name: str = None) -> Dict[str, Any]: if dataset_name is None: dataset_name = self.config.get("name", "CIFAR-10") - - return self.loader.get_dataset_info(dataset_name) \ No newline at end of file + return self.loader.get_dataset_info(dataset_name) diff --git a/libs/fl_core/federated/client_manager.py b/libs/fl_core/federated/client_manager.py index bbb6c07..8d546ee 100644 --- a/libs/fl_core/federated/client_manager.py +++ b/libs/fl_core/federated/client_manager.py @@ -11,8 +11,10 @@ from .client import FederatedClient -from fl_core.compression.sparsification import GlobalTopKSparsifier +from fl_core.compression.sparsification import CompressionStrategy, GlobalTopKSparsifier +from fl_core.privacy.differential_privacy import DifferentialPrivacyManager from fl_core.privacy.encryption import CKKSManager +from fl_core.privacy.secure_aggregation import SecureAggregationMasker class ClientManager: @@ -22,8 +24,10 @@ def __init__(self, selection_strategy: str = "random", max_workers: Optional[int] = None, device: torch.device = None, - sparsifier: Optional[GlobalTopKSparsifier] = None, - ckks_manager: Optional[CKKSManager] = None): + sparsifier: Optional[CompressionStrategy] = None, + ckks_manager: Optional[CKKSManager] = None, + dp_manager: Optional[DifferentialPrivacyManager] = None, + secure_aggregation: Optional[SecureAggregationMasker] = None): # if not clients: # raise ValueError("客户端列表不能为空") @@ -47,6 +51,8 @@ def __init__(self, self.sparsifier = sparsifier self.ckks_manager = ckks_manager + self.dp_manager = dp_manager + self.secure_aggregation = secure_aggregation self.training_stats = { 'total_rounds': 0, @@ -353,7 +359,7 @@ def evaluate_single_client(client: FederatedClient) -> Dict[str, Any]: def get_client_models(self, selected_clients: List[FederatedClient], global_model_params: Optional[Dict[str, torch.Tensor]] = None) -> List[Dict[str, torch.Tensor]]: - client_models = [] + delta_models = [] for client in selected_clients: try: @@ -370,6 +376,9 @@ def get_client_models(self, selected_clients: List[FederatedClient], else: delta = {k: v.to(self.device) for k, v in model_params.items()} + if self.dp_manager: + delta = self.dp_manager.apply(delta) + payload = delta if self.ckks_manager: @@ -379,13 +388,16 @@ def get_client_models(self, selected_clients: List[FederatedClient], # 稀疏化 payload = self.sparsifier.sparsify(delta) - client_models.append(payload) + delta_models.append(payload) except Exception as e: self.logger.error(f"获取客户端 {client.client_id} 模型参数失败: {str(e)}") - client_models.append({}) + delta_models.append({}) - return client_models + if self.secure_aggregation: + delta_models = self.secure_aggregation.mask_models(delta_models) + + return delta_models def broadcast_model_to_clients(self, selected_clients: List[FederatedClient], @@ -459,4 +471,4 @@ def __str__(self) -> str: f"max_workers={self.max_workers})") def __repr__(self) -> str: - return self.__str__() \ No newline at end of file + return self.__str__() diff --git a/libs/fl_core/models/__init__.py b/libs/fl_core/models/__init__.py index ca9209e..36cd1e7 100644 --- a/libs/fl_core/models/__init__.py +++ b/libs/fl_core/models/__init__.py @@ -3,6 +3,15 @@ from .lenet import LeNet, LeNetCIFAR from .cnn import SimpleCNN, DeepCNN, CNNMNIST +from .basic import ( + MclrLogistic, + DNN, + FedAvgCNN, + TextLogistic, + TextDNN, + TextCNN, + CharLSTM, +) from .resnet import ( ResNet18, ResNet34, ResNet50, ResNet18CIFAR, ResNet34CIFAR, @@ -12,9 +21,11 @@ __all__ = [ 'LeNet', 'LeNetCIFAR', + 'MclrLogistic', 'DNN', 'FedAvgCNN', + 'TextLogistic', 'TextDNN', 'TextCNN', 'CharLSTM', 'SimpleCNN', 'DeepCNN', 'CNNMNIST', 'ResNet18', 'ResNet34', 'ResNet50', 'ResNet18CIFAR', 'ResNet34CIFAR', 'ResNet18MNIST', 'ModelManager' -] \ No newline at end of file +] diff --git a/libs/fl_core/models/basic.py b/libs/fl_core/models/basic.py new file mode 100644 index 0000000..71304a5 --- /dev/null +++ b/libs/fl_core/models/basic.py @@ -0,0 +1,122 @@ +from __future__ import annotations + +import torch +import torch.nn as nn +import torch.nn.functional as F + + +def _flatten_dim(input_shape: tuple[int, ...] | list[int] | int | None, input_channels: int = 1) -> int: + if isinstance(input_shape, int): + return input_shape + if input_shape: + total = 1 + for dim in input_shape: + total *= int(dim) + return total + return int(input_channels) * 28 * 28 + + +class MclrLogistic(nn.Module): + def __init__( + self, + num_classes: int = 10, + input_channels: int = 1, + input_shape: tuple[int, ...] | list[int] | int | None = None, + input_dim: int | None = None, + ): + super().__init__() + self.input_dim = int(input_dim or _flatten_dim(input_shape, input_channels)) + self.linear = nn.Linear(self.input_dim, num_classes) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + return self.linear(x.view(x.size(0), -1).float()) + + +class DNN(nn.Module): + def __init__( + self, + num_classes: int = 10, + input_channels: int = 1, + input_shape: tuple[int, ...] | list[int] | int | None = None, + input_dim: int | None = None, + hidden_dim: int = 100, + ): + super().__init__() + self.input_dim = int(input_dim or _flatten_dim(input_shape, input_channels)) + self.fc1 = nn.Linear(self.input_dim, hidden_dim) + self.fc2 = nn.Linear(hidden_dim, num_classes) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + x = x.view(x.size(0), -1).float() + x = F.relu(self.fc1(x)) + return self.fc2(x) + + +class FedAvgCNN(nn.Module): + def __init__(self, num_classes: int = 10, input_channels: int = 3): + super().__init__() + self.conv1 = nn.Conv2d(input_channels, 32, kernel_size=5, padding=2) + self.conv2 = nn.Conv2d(32, 64, kernel_size=5, padding=2) + self.pool = nn.MaxPool2d(2) + self.adaptive_pool = nn.AdaptiveAvgPool2d((4, 4)) + self.fc1 = nn.Linear(64 * 4 * 4, 512) + self.fc2 = nn.Linear(512, num_classes) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + x = self.pool(F.relu(self.conv1(x.float()))) + x = self.pool(F.relu(self.conv2(x))) + x = self.adaptive_pool(x) + x = x.view(x.size(0), -1) + x = F.relu(self.fc1(x)) + return self.fc2(x) + + +class TextLogistic(MclrLogistic): + pass + + +class TextDNN(DNN): + pass + + +class TextCNN(nn.Module): + def __init__( + self, + num_classes: int = 4, + input_dim: int = 5000, + hidden_dim: int = 128, + ): + super().__init__() + self.input_dim = int(input_dim) + self.conv1 = nn.Conv1d(1, 64, kernel_size=5, padding=2) + self.conv2 = nn.Conv1d(64, hidden_dim, kernel_size=5, padding=2) + self.pool = nn.AdaptiveMaxPool1d(1) + self.fc = nn.Linear(hidden_dim, num_classes) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + x = x.view(x.size(0), 1, -1).float() + x = F.relu(self.conv1(x)) + x = F.relu(self.conv2(x)) + x = self.pool(x).squeeze(-1) + return self.fc(x) + + +class CharLSTM(nn.Module): + def __init__( + self, + num_classes: int = 128, + vocab_size: int = 128, + embed_dim: int = 32, + hidden_dim: int = 128, + ): + super().__init__() + self.vocab_size = int(vocab_size) + self.embedding = nn.Embedding(self.vocab_size, embed_dim) + self.lstm = nn.LSTM(embed_dim, hidden_dim, batch_first=True) + self.fc = nn.Linear(hidden_dim, num_classes) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + token_ids = x.long().clamp(min=0, max=self.vocab_size - 1) + embedded = self.embedding(token_ids) + _, (hidden, _) = self.lstm(embedded) + return self.fc(hidden[-1]) diff --git a/libs/fl_core/models/model_manager.py b/libs/fl_core/models/model_manager.py index c00579e..bdf3447 100644 --- a/libs/fl_core/models/model_manager.py +++ b/libs/fl_core/models/model_manager.py @@ -1,105 +1,237 @@ -import torch -import torch.nn as nn +from __future__ import annotations + import copy import pickle -from typing import Dict, Any, Optional, Union -from collections import OrderedDict +from typing import Any, Dict, Optional + +import torch +import torch.nn as nn +from fl_core.simulation_registry import canonical_model_name, list_model_names + +from .basic import ( + CharLSTM, + DNN, + FedAvgCNN, + MclrLogistic, + TextCNN, + TextDNN, + TextLogistic, +) +from .cnn import CNNMNIST, DeepCNN, SimpleCNN from .lenet import LeNet, LeNetCIFAR -from .cnn import SimpleCNN, DeepCNN, CNNMNIST from .resnet import ( - ResNet18, ResNet34, ResNet50, - ResNet18CIFAR, ResNet34CIFAR, - ResNet18MNIST + ResNet18, + ResNet18CIFAR, + ResNet18MNIST, + ResNet34, + ResNet34CIFAR, + ResNet50, ) class ModelManager: - def __init__(self, config: Dict[str, Any]): self.config = config - self.device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') - - self.model_registry = { - 'lenet': LeNet, - 'lenet_cifar': LeNetCIFAR, - - 'cnn': SimpleCNN, - 'simple_cnn': SimpleCNN, - 'deep_cnn': DeepCNN, - 'cnn_mnist': CNNMNIST, - - 'resnet18': ResNet18, - 'resnet34': ResNet34, - 'resnet50': ResNet50, - 'resnet18_cifar': ResNet18CIFAR, - 'resnet34_cifar': ResNet34CIFAR, - 'resnet18_mnist': ResNet18MNIST, + self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu") + + self.model_registry: dict[str, Any] = { + "mclr_logistic": MclrLogistic, + "mclr": MclrLogistic, + "logistic": MclrLogistic, + "logistic_regression": MclrLogistic, + "dnn": DNN, + "mlp": DNN, + "fedavgcnn": FedAvgCNN, + "fedavg_cnn": FedAvgCNN, + "lenet": LeNet, + "lenet_cifar": LeNetCIFAR, + "cnn": SimpleCNN, + "simple_cnn": SimpleCNN, + "deep_cnn": DeepCNN, + "cnn_mnist": CNNMNIST, + "resnet18": ResNet18, + "resnet34": ResNet34, + "resnet50": ResNet50, + "resnet18_cifar": ResNet18CIFAR, + "resnet34_cifar": ResNet34CIFAR, + "resnet18_mnist": ResNet18MNIST, + "alexnet": "torchvision_alexnet", + "mobilenet": "torchvision_mobilenet_v2", + "mobilenet_v2": "torchvision_mobilenet_v2", + "mobilenetv2": "torchvision_mobilenet_v2", + "googlenet": "torchvision_googlenet", + "google_net": "torchvision_googlenet", + "textlogistic": TextLogistic, + "text_logistic": TextLogistic, + "textdnn": TextDNN, + "text_dnn": TextDNN, + "textcnn": TextCNN, + "text_cnn": TextCNN, + "charlstm": CharLSTM, + "char_lstm": CharLSTM, + "char_rnn": CharLSTM, } - - def create_model(self, model_name: str, input_shape: tuple, num_classes: int) -> nn.Module: - model_name_lower = model_name.lower() - - if model_name_lower not in self.model_registry: - raise ValueError(f"不支持的模型类型: {model_name}. 支持的模型: {list(self.model_registry.keys())}") - - model_class = self.model_registry[model_name_lower] - input_channels = input_shape[0] if len(input_shape) == 3 else 1 - - if model_name_lower == 'lenet': - if input_shape[1:] == (32, 32): # CIFAR数据集 + + def create_model( + self, + model_name: str, + input_shape: tuple, + num_classes: int, + dataset_info: Optional[Dict[str, Any]] = None, + ) -> nn.Module: + canonical_name = canonical_model_name(model_name) + registry_key = canonical_name.lower().replace("-", "_") + + if registry_key not in self.model_registry: + raise ValueError( + f"Unsupported model type: {model_name}. " + f"Supported models: {list(self.model_registry.keys())}" + ) + + model_class = self.model_registry[registry_key] + dataset_info = dataset_info or {} + input_shape_tuple = tuple(input_shape or ()) + input_channels = input_shape_tuple[0] if len(input_shape_tuple) == 3 else 1 + input_dim = self._flatten_dim(input_shape_tuple) + modality = dataset_info.get("modality", "image") + + if registry_key == "lenet": + if len(input_shape_tuple) == 3 and input_shape_tuple[1:] == (32, 32): model = LeNetCIFAR(num_classes=num_classes, input_channels=input_channels) - else: # MNIST数据集 + else: model = LeNet(num_classes=num_classes, input_channels=input_channels) - elif model_name_lower in ['cnn', 'simple_cnn']: - if input_shape[1:] == (28, 28): # MNIST数据集 + elif registry_key == "fedavgcnn": + model = FedAvgCNN(num_classes=num_classes, input_channels=input_channels) + elif registry_key in {"cnn", "simple_cnn"}: + if len(input_shape_tuple) == 3 and input_shape_tuple[1:] == (28, 28): model = CNNMNIST(num_classes=num_classes, input_channels=input_channels) - else: # CIFAR数据集 + else: model = SimpleCNN(num_classes=num_classes, input_channels=input_channels) - elif model_name_lower == 'resnet18': - if input_shape[1:] == (32, 32): # CIFAR数据集 + elif registry_key == "resnet18": + if len(input_shape_tuple) == 3 and input_shape_tuple[1:] == (32, 32): model = ResNet18CIFAR(num_classes=num_classes, input_channels=input_channels) - elif input_shape[1:] == (28, 28): # MNIST数据集 + elif len(input_shape_tuple) == 3 and input_shape_tuple[1:] == (28, 28): model = ResNet18MNIST(num_classes=num_classes, input_channels=input_channels) - else: # 标准ImageNet尺寸 + else: model = ResNet18(num_classes=num_classes, input_channels=input_channels) + elif registry_key in {"mclr_logistic", "mclr", "logistic", "logistic_regression"}: + model = MclrLogistic(num_classes=num_classes, input_shape=input_shape_tuple, input_dim=input_dim) + elif registry_key in {"dnn", "mlp"}: + model = DNN( + num_classes=num_classes, + input_shape=input_shape_tuple, + input_dim=input_dim, + hidden_dim=int(self.config.get("hidden_dim", 100)), + ) + elif registry_key in {"textlogistic", "text_logistic"}: + model = TextLogistic( + num_classes=num_classes, + input_dim=input_dim or int(dataset_info.get("feature_dim", 5000)), + ) + elif registry_key in {"textdnn", "text_dnn"}: + model = TextDNN( + num_classes=num_classes, + input_dim=input_dim or int(dataset_info.get("feature_dim", 5000)), + hidden_dim=int(self.config.get("hidden_dim", 100)), + ) + elif registry_key in {"textcnn", "text_cnn"}: + model = TextCNN( + num_classes=num_classes, + input_dim=input_dim or int(dataset_info.get("feature_dim", 5000)), + ) + elif registry_key in {"charlstm", "char_lstm", "char_rnn"}: + model = CharLSTM( + num_classes=num_classes, + vocab_size=int(dataset_info.get("vocab_size", num_classes)), + hidden_dim=int(self.config.get("hidden_dim", 128)), + ) + elif isinstance(model_class, str) and model_class.startswith("torchvision_"): + if modality != "image": + raise ValueError(f"{canonical_name} only supports image datasets") + model = self._create_torchvision_model(model_class, num_classes, input_channels) else: - # 直接使用指定的模型类 - model = model_class(num_classes=num_classes, input_channels=input_channels) - + try: + model = model_class(num_classes=num_classes, input_channels=input_channels) + except TypeError: + model = model_class(num_classes=num_classes) + model.to(self.device) return model - + + @staticmethod + def _flatten_dim(input_shape: tuple) -> Optional[int]: + if not input_shape: + return None + total = 1 + for dim in input_shape: + total *= int(dim) + return total + + def _create_torchvision_model(self, model_key: str, num_classes: int, input_channels: int) -> nn.Module: + import torchvision.models as tv_models + + if model_key == "torchvision_alexnet": + model = tv_models.alexnet(weights=None, num_classes=num_classes) + if input_channels != 3: + model.features[0] = nn.Conv2d(input_channels, 64, kernel_size=11, stride=4, padding=2) + return model + + if model_key == "torchvision_mobilenet_v2": + model = tv_models.mobilenet_v2(weights=None, num_classes=num_classes) + if input_channels != 3: + model.features[0][0] = nn.Conv2d( + input_channels, + 32, + kernel_size=3, + stride=2, + padding=1, + bias=False, + ) + return model + + if model_key == "torchvision_googlenet": + model = tv_models.googlenet(weights=None, aux_logits=False, num_classes=num_classes) + if input_channels != 3: + model.conv1.conv = nn.Conv2d( + input_channels, + 64, + kernel_size=7, + stride=2, + padding=3, + bias=False, + ) + return model + + raise ValueError(f"Unsupported torchvision model key: {model_key}") + def get_model_parameters(self, model: nn.Module) -> Dict[str, torch.Tensor]: return {name: param.clone().detach() for name, param in model.named_parameters()} - + def get_model_state_dict(self, model: nn.Module) -> Dict[str, torch.Tensor]: return {name: param.clone().detach() for name, param in model.state_dict().items()} - + def set_model_parameters(self, model: nn.Module, parameters: Dict[str, torch.Tensor]) -> None: model_dict = model.state_dict() - - # 验证参数键是否匹配 + param_keys = set(parameters.keys()) model_keys = set(model_dict.keys()) - if param_keys != model_keys: missing_keys = model_keys - param_keys unexpected_keys = param_keys - model_keys - if missing_keys: - print(f"警告: 缺少参数键: {missing_keys}") + print(f"Warning: missing parameter keys: {missing_keys}") if unexpected_keys: - print(f"警告: 意外的参数键: {unexpected_keys}") - + print(f"Warning: unexpected parameter keys: {unexpected_keys}") + for name, param in parameters.items(): if name in model_dict: model_dict[name].copy_(param.to(self.device)) - + def set_model_state_dict(self, model: nn.Module, state_dict: Dict[str, torch.Tensor]) -> None: device_state_dict = {name: param.to(self.device) for name, param in state_dict.items()} model.load_state_dict(device_state_dict) - + def get_model_gradients(self, model: nn.Module) -> Dict[str, torch.Tensor]: gradients = {} for name, param in model.named_parameters(): @@ -108,7 +240,7 @@ def get_model_gradients(self, model: nn.Module) -> Dict[str, torch.Tensor]: else: gradients[name] = torch.zeros_like(param) return gradients - + def set_model_gradients(self, model: nn.Module, gradients: Dict[str, torch.Tensor]) -> None: for name, param in model.named_parameters(): if name in gradients: @@ -116,68 +248,68 @@ def set_model_gradients(self, model: nn.Module, gradients: Dict[str, torch.Tenso param.grad = gradients[name].clone().to(self.device) else: param.grad.copy_(gradients[name].to(self.device)) - + def serialize_parameters(self, parameters: Dict[str, torch.Tensor]) -> bytes: cpu_parameters = {name: param.cpu() for name, param in parameters.items()} return pickle.dumps(cpu_parameters) - + def deserialize_parameters(self, data: bytes) -> Dict[str, torch.Tensor]: parameters = pickle.loads(data) return {name: param.to(self.device) for name, param in parameters.items()} - + def save_model(self, model: nn.Module, filepath: str) -> None: - torch.save({ - 'model_state_dict': model.state_dict(), - 'model_config': self.config, - 'model_class': model.__class__.__name__ - }, filepath) - + torch.save( + { + "model_state_dict": model.state_dict(), + "model_config": self.config, + "model_class": model.__class__.__name__, + }, + filepath, + ) + def load_model(self, filepath: str, model_name: str, input_shape: tuple, num_classes: int) -> nn.Module: checkpoint = torch.load(filepath, map_location=self.device) - model = self.create_model(model_name, input_shape, num_classes) - - model.load_state_dict(checkpoint['model_state_dict']) - + model.load_state_dict(checkpoint["model_state_dict"]) return model - + def clone_model(self, model: nn.Module) -> nn.Module: cloned_model = copy.deepcopy(model) cloned_model.to(self.device) return cloned_model - + def get_model_size(self, model: nn.Module) -> int: return sum(p.numel() for p in model.parameters()) - + def get_trainable_parameters(self, model: nn.Module) -> int: return sum(p.numel() for p in model.parameters() if p.requires_grad) - + def freeze_layers(self, model: nn.Module, layer_names: list) -> None: for name, param in model.named_parameters(): if any(layer_name in name for layer_name in layer_names): param.requires_grad = False - + def unfreeze_layers(self, model: nn.Module, layer_names: list) -> None: for name, param in model.named_parameters(): if any(layer_name in name for layer_name in layer_names): param.requires_grad = True - + def get_supported_models(self) -> list: - return list(self.model_registry.keys()) - + return list_model_names() + def model_summary(self, model: nn.Module, input_shape: tuple) -> str: total_params = self.get_model_size(model) trainable_params = self.get_trainable_parameters(model) - + summary = f""" -模型摘要: +Model summary: ======== -模型类型: {model.__class__.__name__} -输入形状: {input_shape} -总参数数: {total_params:,} -可训练参数: {trainable_params:,} -设备: {self.device} +Model type: {model.__class__.__name__} +Input shape: {input_shape} +Total parameters: {total_params:,} +Trainable parameters: {trainable_params:,} +Device: {self.device} ======== """ - - return summary.strip() \ No newline at end of file + + return summary.strip() diff --git a/libs/fl_core/models/resnet.py b/libs/fl_core/models/resnet.py index cfc8c2d..6445c4a 100644 --- a/libs/fl_core/models/resnet.py +++ b/libs/fl_core/models/resnet.py @@ -88,7 +88,7 @@ def forward(self, x): out = self.layer2(out) out = self.layer3(out) out = self.layer4(out) - out = F.avg_pool2d(out, 4) + out = F.adaptive_avg_pool2d(out, 1) out = out.view(out.size(0), -1) out = self.linear(out) return out @@ -138,7 +138,7 @@ def forward(self, x): out = self.layer1(out) out = self.layer2(out) out = self.layer3(out) - out = F.avg_pool2d(out, 8) + out = F.adaptive_avg_pool2d(out, 1) out = out.view(out.size(0), -1) out = self.linear(out) return out @@ -185,7 +185,7 @@ def forward(self, x): out = self.layer1(out) out = self.layer2(out) out = self.layer3(out) - out = F.avg_pool2d(out, 7) + out = F.adaptive_avg_pool2d(out, 1) out = out.view(out.size(0), -1) out = self.linear(out) return out @@ -195,4 +195,4 @@ def get_feature_dim(self): def ResNet18MNIST(num_classes=10, input_channels=1): - return ResNetMNIST(BasicBlock, [2, 2, 2], num_classes, input_channels) \ No newline at end of file + return ResNetMNIST(BasicBlock, [2, 2, 2], num_classes, input_channels) diff --git a/libs/fl_core/privacy/__init__.py b/libs/fl_core/privacy/__init__.py index e69de29..ee404a8 100644 --- a/libs/fl_core/privacy/__init__.py +++ b/libs/fl_core/privacy/__init__.py @@ -0,0 +1,9 @@ +from .differential_privacy import DifferentialPrivacyManager +from .encryption import CKKSManager +from .secure_aggregation import SecureAggregationMasker + +__all__ = [ + "CKKSManager", + "DifferentialPrivacyManager", + "SecureAggregationMasker", +] diff --git a/libs/fl_core/privacy/differential_privacy.py b/libs/fl_core/privacy/differential_privacy.py new file mode 100644 index 0000000..19d81f2 --- /dev/null +++ b/libs/fl_core/privacy/differential_privacy.py @@ -0,0 +1,54 @@ +from __future__ import annotations + +import logging +from typing import Dict + +import torch + + +def _is_trainable_float(name: str, tensor: torch.Tensor) -> bool: + return ( + tensor.is_floating_point() + and tensor.dim() > 0 + and "num_batches_tracked" not in name + and "running_" not in name + ) + + +class DifferentialPrivacyManager: + def __init__(self, clipping_norm: float = 1.0, noise_multiplier: float = 0.0): + self.clipping_norm = float(clipping_norm) + self.noise_multiplier = float(noise_multiplier) + self.logger = logging.getLogger("DifferentialPrivacyManager") + + def apply(self, update_dict: Dict[str, torch.Tensor]) -> Dict[str, torch.Tensor]: + trainable = { + name: tensor + for name, tensor in update_dict.items() + if _is_trainable_float(name, tensor) + } + if not trainable: + return update_dict + + total_norm = torch.sqrt( + sum(torch.sum(tensor.detach() ** 2) for tensor in trainable.values()) + ) + clip_factor = min(1.0, self.clipping_norm / (total_norm.item() + 1e-12)) + noise_std = self.noise_multiplier * self.clipping_norm + + protected = {} + for name, tensor in update_dict.items(): + if not _is_trainable_float(name, tensor): + protected[name] = tensor + continue + value = tensor * clip_factor + if noise_std > 0: + value = value + torch.normal( + mean=0.0, + std=noise_std, + size=tensor.shape, + device=tensor.device, + dtype=tensor.dtype, + ) + protected[name] = value + return protected diff --git a/libs/fl_core/privacy/secure_aggregation.py b/libs/fl_core/privacy/secure_aggregation.py new file mode 100644 index 0000000..4d4e501 --- /dev/null +++ b/libs/fl_core/privacy/secure_aggregation.py @@ -0,0 +1,54 @@ +from __future__ import annotations + +import logging +from typing import Dict, List + +import torch + + +def _is_trainable_float(name: str, tensor: torch.Tensor) -> bool: + return ( + tensor.is_floating_point() + and tensor.dim() > 0 + and "num_batches_tracked" not in name + and "running_" not in name + ) + + +class SecureAggregationMasker: + """ + Simulation-only additive masking. + + Masks are generated so their unweighted sum is zero. This is suitable for + equal-weight FedAvg/simple average in local simulation; it is intentionally + not a production secure aggregation protocol. + """ + + def __init__(self, mask_std: float = 1.0): + self.mask_std = float(mask_std) + self.logger = logging.getLogger("SecureAggregationMasker") + + def mask_models(self, updates: List[Dict[str, torch.Tensor]]) -> List[Dict[str, torch.Tensor]]: + if len(updates) <= 1 or self.mask_std <= 0: + return updates + + masked = [{name: tensor.clone() for name, tensor in update.items()} for update in updates] + reference = updates[0] + + for name, tensor in reference.items(): + if not _is_trainable_float(name, tensor): + continue + running_sum = torch.zeros_like(tensor) + for idx in range(len(updates) - 1): + mask = torch.normal( + mean=0.0, + std=self.mask_std, + size=tensor.shape, + device=tensor.device, + dtype=tensor.dtype, + ) + masked[idx][name] = masked[idx][name] + mask + running_sum = running_sum + mask + masked[-1][name] = masked[-1][name] - running_sum + + return masked diff --git a/libs/fl_core/simulation_registry.py b/libs/fl_core/simulation_registry.py new file mode 100644 index 0000000..1f71563 --- /dev/null +++ b/libs/fl_core/simulation_registry.py @@ -0,0 +1,528 @@ +from __future__ import annotations + +from dataclasses import dataclass, field +from typing import Any + + +@dataclass(frozen=True) +class DatasetSpec: + name: str + aliases: tuple[str, ...] + modality: str + task_type: str + input_shape: tuple[int, ...] + num_classes: int + default_model: str + compatible_models: tuple[str, ...] + supported_splits: tuple[str, ...] = ("iid", "non_iid") + loader: str = "torchvision" + requires_local_data: bool = False + metadata: dict[str, Any] = field(default_factory=dict) + + +@dataclass(frozen=True) +class ModelSpec: + name: str + aliases: tuple[str, ...] + modalities: tuple[str, ...] + task_types: tuple[str, ...] + description: str = "" + + +LINEAR_AGGREGATIONS = ("fedavg", "weighted_avg", "simple_avg") +EQUAL_WEIGHT_AGGREGATIONS = ("fedavg", "simple_avg") +EXECUTABLE_AGGREGATIONS = (*LINEAR_AGGREGATIONS,) +COMPRESSION_METHODS = ("global_topk", "random_k", "threshold", "sign", "quant_int8") + + +MODEL_SPECS: tuple[ModelSpec, ...] = ( + ModelSpec( + name="Mclr_Logistic", + aliases=("mclr", "logistic", "logistic_regression", "MCLR"), + modalities=("image", "text"), + task_types=("classification",), + description="Multiclass logistic regression over flattened features.", + ), + ModelSpec( + name="LeNet", + aliases=("lenet",), + modalities=("image",), + task_types=("classification",), + ), + ModelSpec( + name="DNN", + aliases=("dnn", "mlp"), + modalities=("image", "text"), + task_types=("classification",), + description="Two-layer dense network with configurable flattened input.", + ), + ModelSpec( + name="FedAvgCNN", + aliases=("cnn", "simple_cnn", "fedavg_cnn"), + modalities=("image",), + task_types=("classification",), + ), + ModelSpec( + name="ResNet18", + aliases=("resnet", "resnet18"), + modalities=("image",), + task_types=("classification",), + ), + ModelSpec( + name="ResNet34", + aliases=("resnet34",), + modalities=("image",), + task_types=("classification",), + ), + ModelSpec( + name="ResNet50", + aliases=("resnet50",), + modalities=("image",), + task_types=("classification",), + ), + ModelSpec( + name="AlexNet", + aliases=("alexnet",), + modalities=("image",), + task_types=("classification",), + ), + ModelSpec( + name="MobileNet", + aliases=("mobilenet", "mobilenet_v2", "mobilenetv2"), + modalities=("image",), + task_types=("classification",), + ), + ModelSpec( + name="GoogleNet", + aliases=("googlenet", "google_net"), + modalities=("image",), + task_types=("classification",), + ), + ModelSpec( + name="TextLogistic", + aliases=("text_logistic",), + modalities=("text",), + task_types=("classification",), + ), + ModelSpec( + name="TextDNN", + aliases=("text_dnn",), + modalities=("text",), + task_types=("classification",), + ), + ModelSpec( + name="TextCNN", + aliases=("text_cnn",), + modalities=("text",), + task_types=("classification",), + ), + ModelSpec( + name="CharLSTM", + aliases=("char_lstm", "lstm", "char_rnn"), + modalities=("sequence",), + task_types=("next_char",), + ), +) + + +DIGIT_MODELS = ("Mclr_Logistic", "LeNet", "DNN") +SMALL_IMAGE_MODELS = ( + "Mclr_Logistic", + "FedAvgCNN", + "DNN", + "ResNet18", + "AlexNet", + "MobileNet", + "GoogleNet", +) +GENERAL_IMAGE_MODELS = ( + "FedAvgCNN", + "DNN", + "ResNet18", + "ResNet34", + "AlexNet", + "MobileNet", + "GoogleNet", +) +TEXT_MODELS = ("TextLogistic", "TextDNN", "TextCNN") +SEQUENCE_MODELS = ("CharLSTM",) + + +DATASET_SPECS: tuple[DatasetSpec, ...] = ( + DatasetSpec( + name="MNIST", + aliases=("mnist",), + modality="image", + task_type="classification", + input_shape=(1, 28, 28), + num_classes=10, + default_model="LeNet", + compatible_models=DIGIT_MODELS, + ), + DatasetSpec( + name="EMNIST", + aliases=("emnist",), + modality="image", + task_type="classification", + input_shape=(1, 28, 28), + num_classes=47, + default_model="LeNet", + compatible_models=DIGIT_MODELS, + metadata={"split": "balanced"}, + ), + DatasetSpec( + name="FEMNIST", + aliases=("femnist",), + modality="image", + task_type="classification", + input_shape=(1, 28, 28), + num_classes=62, + default_model="LeNet", + compatible_models=DIGIT_MODELS, + loader="leaf_femnist", + requires_local_data=True, + ), + DatasetSpec( + name="Fashion-MNIST", + aliases=("fashion-mnist", "fashion_mnist", "fashionmnist"), + modality="image", + task_type="classification", + input_shape=(1, 28, 28), + num_classes=10, + default_model="LeNet", + compatible_models=DIGIT_MODELS, + ), + DatasetSpec( + name="CIFAR-10", + aliases=("cifar10", "cifar-10", "Cifar10"), + modality="image", + task_type="classification", + input_shape=(3, 32, 32), + num_classes=10, + default_model="FedAvgCNN", + compatible_models=SMALL_IMAGE_MODELS, + ), + DatasetSpec( + name="CIFAR-100", + aliases=("cifar100", "cifar-100", "Cifar100"), + modality="image", + task_type="classification", + input_shape=(3, 32, 32), + num_classes=100, + default_model="FedAvgCNN", + compatible_models=SMALL_IMAGE_MODELS, + ), + DatasetSpec( + name="AG News", + aliases=("ag_news", "agnews"), + modality="text", + task_type="classification", + input_shape=(5000,), + num_classes=4, + default_model="TextDNN", + compatible_models=TEXT_MODELS, + loader="csv_text", + requires_local_data=True, + metadata={"feature_dim": 5000}, + ), + DatasetSpec( + name="Sogou News", + aliases=("sogou_news", "sogou"), + modality="text", + task_type="classification", + input_shape=(8000,), + num_classes=5, + default_model="TextDNN", + compatible_models=TEXT_MODELS, + loader="csv_text", + requires_local_data=True, + metadata={"feature_dim": 8000}, + ), + DatasetSpec( + name="Tiny-ImageNet", + aliases=("tiny_imagenet", "tinyimagenet", "tiny-imagenet"), + modality="image", + task_type="classification", + input_shape=(3, 32, 32), + num_classes=200, + default_model="FedAvgCNN", + compatible_models=SMALL_IMAGE_MODELS, + loader="tiny_imagenet", + requires_local_data=True, + ), + DatasetSpec( + name="Country211", + aliases=("country211",), + modality="image", + task_type="classification", + input_shape=(3, 64, 64), + num_classes=211, + default_model="MobileNet", + compatible_models=GENERAL_IMAGE_MODELS, + ), + DatasetSpec( + name="Flowers102", + aliases=("flowers102", "flowers-102"), + modality="image", + task_type="classification", + input_shape=(3, 64, 64), + num_classes=102, + default_model="MobileNet", + compatible_models=GENERAL_IMAGE_MODELS, + ), + DatasetSpec( + name="GTSRB", + aliases=("gtsrb",), + modality="image", + task_type="classification", + input_shape=(3, 64, 64), + num_classes=43, + default_model="MobileNet", + compatible_models=GENERAL_IMAGE_MODELS, + ), + DatasetSpec( + name="Shakespeare", + aliases=("shakespeare",), + modality="sequence", + task_type="next_char", + input_shape=(80,), + num_classes=128, + default_model="CharLSTM", + compatible_models=SEQUENCE_MODELS, + loader="shakespeare", + requires_local_data=True, + metadata={"sequence_length": 80, "vocab_size": 128}, + ), + DatasetSpec( + name="Stanford Cars", + aliases=("stanford_cars", "stanford-cars", "cars"), + modality="image", + task_type="classification", + input_shape=(3, 64, 64), + num_classes=196, + default_model="MobileNet", + compatible_models=GENERAL_IMAGE_MODELS, + ), + DatasetSpec( + name="COVIDx", + aliases=("covidx", "covid-x"), + modality="image", + task_type="classification", + input_shape=(3, 64, 64), + num_classes=3, + default_model="MobileNet", + compatible_models=GENERAL_IMAGE_MODELS, + loader="local_image_folder", + requires_local_data=True, + ), + DatasetSpec( + name="Kvasir", + aliases=("kvasir",), + modality="image", + task_type="classification", + input_shape=(3, 64, 64), + num_classes=8, + default_model="MobileNet", + compatible_models=GENERAL_IMAGE_MODELS, + loader="local_image_folder", + requires_local_data=True, + ), +) + + +def _normalize_key(value: str) -> str: + return value.strip().lower().replace("_", "-").replace(" ", "-") + + +DATASET_BY_NAME: dict[str, DatasetSpec] = { + _normalize_key(spec.name): spec for spec in DATASET_SPECS +} +for _spec in DATASET_SPECS: + for _alias in _spec.aliases: + DATASET_BY_NAME[_normalize_key(_alias)] = _spec + + +MODEL_BY_NAME: dict[str, ModelSpec] = { + _normalize_key(spec.name): spec for spec in MODEL_SPECS +} +for _spec in MODEL_SPECS: + for _alias in _spec.aliases: + MODEL_BY_NAME[_normalize_key(_alias)] = _spec + + +def canonical_dataset_name(name: str) -> str: + spec = get_dataset_spec(name) + if spec is None: + return name + return spec.name + + +def canonical_model_name(name: str) -> str: + spec = get_model_spec(name) + if spec is None: + return name + return spec.name + + +def get_dataset_spec(name: str | None) -> DatasetSpec | None: + if not name: + return None + return DATASET_BY_NAME.get(_normalize_key(str(name))) + + +def get_model_spec(name: str | None) -> ModelSpec | None: + if not name: + return None + return MODEL_BY_NAME.get(_normalize_key(str(name))) + + +def list_dataset_names() -> list[str]: + return [spec.name for spec in DATASET_SPECS] + + +def list_model_names() -> list[str]: + return [spec.name for spec in MODEL_SPECS] + + +def is_model_compatible_with_dataset(model_name: str, dataset_name: str) -> bool: + dataset = get_dataset_spec(dataset_name) + model = get_model_spec(model_name) + if dataset is None or model is None: + return False + return model.name in dataset.compatible_models + + +def get_compatible_models(dataset_name: str) -> list[str]: + dataset = get_dataset_spec(dataset_name) + if dataset is None: + return [] + return list(dataset.compatible_models) + + +def resolve_model_name_for_dataset(model_name: str | None, dataset_name: str) -> str: + dataset = get_dataset_spec(dataset_name) + if model_name is None or not str(model_name).strip() or str(model_name).strip().lower() == "auto": + return dataset.default_model if dataset is not None else "FedAvgCNN" + return canonical_model_name(str(model_name)) + + +def validate_training_combination(config: dict[str, Any]) -> list[str]: + errors: list[str] = [] + dataset_cfg = config.get("dataset") if isinstance(config.get("dataset"), dict) else {} + model_cfg = config.get("model") if isinstance(config.get("model"), dict) else {} + federated_cfg = config.get("federated") if isinstance(config.get("federated"), dict) else {} + privacy_cfg = config.get("privacy") if isinstance(config.get("privacy"), dict) else {} + compression_cfg = config.get("compression") if isinstance(config.get("compression"), dict) else {} + + dataset_name = str(dataset_cfg.get("name", "")) + raw_model_name = model_cfg.get("name") + aggregation = str(federated_cfg.get("aggregation", "fedavg")).lower() + + dataset = get_dataset_spec(dataset_name) + model_name = resolve_model_name_for_dataset(raw_model_name, dataset_name) + model = get_model_spec(model_name) + + if dataset is None: + errors.append(f"Unsupported dataset: {dataset_name}") + if model is None: + errors.append(f"Unsupported model: {model_name}") + if dataset is not None and model is not None and model.name not in dataset.compatible_models: + allowed = ", ".join(dataset.compatible_models) + errors.append(f"Model {model.name} is not compatible with dataset {dataset.name}. Allowed models: {allowed}") + + if aggregation not in EXECUTABLE_AGGREGATIONS: + allowed = ", ".join(EXECUTABLE_AGGREGATIONS) + errors.append(f"Aggregation {aggregation} is not executable in simulation. Allowed aggregations: {allowed}") + + he_cfg = privacy_cfg.get("homomorphic_encryption") if isinstance(privacy_cfg, dict) else {} + ckks_enabled = isinstance(he_cfg, dict) and he_cfg.get("enable") is True + dp_cfg = privacy_cfg.get("differential_privacy") if isinstance(privacy_cfg, dict) else {} + dp_enabled = isinstance(dp_cfg, dict) and dp_cfg.get("enable") is True + secure_agg_cfg = privacy_cfg.get("secure_aggregation") if isinstance(privacy_cfg, dict) else {} + secure_agg_enabled = isinstance(secure_agg_cfg, dict) and secure_agg_cfg.get("enable") is True + sparsification_cfg = compression_cfg.get("sparsification") if isinstance(compression_cfg, dict) else {} + sparsification_enabled = isinstance(sparsification_cfg, dict) and sparsification_cfg.get("enable") is True + compression_method = str(sparsification_cfg.get("method", "global_topk")).lower() if isinstance(sparsification_cfg, dict) else "global_topk" + + if ckks_enabled and aggregation not in LINEAR_AGGREGATIONS: + allowed = ", ".join(LINEAR_AGGREGATIONS) + errors.append(f"CKKS is only compatible with linear aggregations: {allowed}") + if ckks_enabled and sparsification_enabled: + errors.append("CKKS homomorphic encryption and sparsification are not compatible yet") + if ckks_enabled and secure_agg_enabled: + errors.append("CKKS homomorphic encryption and secure aggregation masking cannot be enabled together") + if secure_agg_enabled and aggregation not in EQUAL_WEIGHT_AGGREGATIONS: + allowed = ", ".join(EQUAL_WEIGHT_AGGREGATIONS) + errors.append(f"Secure aggregation masking is only compatible with equal-weight aggregations: {allowed}") + if secure_agg_enabled and sparsification_enabled: + errors.append("Secure aggregation masking and compression are not compatible yet") + if sparsification_enabled and aggregation not in LINEAR_AGGREGATIONS: + allowed = ", ".join(LINEAR_AGGREGATIONS) + errors.append(f"Sparsification is only compatible with linear aggregations: {allowed}") + if sparsification_enabled and compression_method not in COMPRESSION_METHODS: + allowed = ", ".join(COMPRESSION_METHODS) + errors.append(f"Unsupported compression method: {compression_method}. Allowed methods: {allowed}") + if dp_enabled: + clipping_norm = dp_cfg.get("clipping_norm", 1.0) + noise_multiplier = dp_cfg.get("noise_multiplier", 0.0) + if isinstance(clipping_norm, bool) or not isinstance(clipping_norm, (int, float)) or clipping_norm <= 0: + errors.append("Differential privacy clipping_norm must be greater than 0") + if isinstance(noise_multiplier, bool) or not isinstance(noise_multiplier, (int, float)) or noise_multiplier < 0: + errors.append("Differential privacy noise_multiplier must be greater than or equal to 0") + + return errors + + +def capabilities_payload() -> dict[str, Any]: + return { + "datasets": list_dataset_names(), + "distributions": ["iid", "non_iid"], + "models": ["Auto", *list_model_names()], + "dataset_model_compatibility": { + spec.name: list(spec.compatible_models) for spec in DATASET_SPECS + }, + "dataset_defaults": { + spec.name: { + "default_model": spec.default_model, + "modality": spec.modality, + "task_type": spec.task_type, + "input_shape": list(spec.input_shape), + "num_classes": spec.num_classes, + "requires_local_data": spec.requires_local_data, + } + for spec in DATASET_SPECS + }, + "aggregations": list(EXECUTABLE_AGGREGATIONS), + "privacy": { + "homomorphic_encryption": { + "methods": ["ckks"], + "compatible_aggregations": list(LINEAR_AGGREGATIONS), + "incompatible_with": ["compression.sparsification", "privacy.secure_aggregation"], + }, + "differential_privacy": { + "methods": ["clip_and_noise"], + "compatible_aggregations": list(EXECUTABLE_AGGREGATIONS), + "compatible_with": [ + "privacy.homomorphic_encryption", + "privacy.secure_aggregation", + "compression.sparsification", + ], + }, + "secure_aggregation": { + "methods": ["additive_masking_simulation"], + "compatible_aggregations": list(EQUAL_WEIGHT_AGGREGATIONS), + "incompatible_with": ["privacy.homomorphic_encryption", "compression.sparsification"], + }, + }, + "compression": { + "sparsification": { + "methods": list(COMPRESSION_METHODS), + "compatible_aggregations": list(LINEAR_AGGREGATIONS), + "incompatible_with": ["privacy.homomorphic_encryption", "privacy.secure_aggregation"], + } + }, + "metrics": { + "global_results": ["rounds", "global_loss", "global_accuracy"], + "client_results": ["train_loss", "train_acc", "test_loss", "test_acc"], + }, + } diff --git a/libs/fl_core/utils/config.py b/libs/fl_core/utils/config.py index b455844..061f433 100644 --- a/libs/fl_core/utils/config.py +++ b/libs/fl_core/utils/config.py @@ -2,6 +2,8 @@ import os from typing import Dict, Any +from fl_core.simulation_registry import validate_training_combination + class ConfigManager: @@ -78,4 +80,8 @@ def validate_config(self) -> bool: if param not in federated_config: raise ValueError(f"联邦学习配置缺少{param}参数") - return True \ No newline at end of file + compatibility_errors = validate_training_combination(self.config) + if compatibility_errors: + raise ValueError("; ".join(compatibility_errors)) + + return True From 913c1b7830c8c62c17623d83d4a82ad1c17ba5ce Mon Sep 17 00:00:00 2001 From: Fennel1 <627235787@qq.com> Date: Thu, 14 May 2026 00:01:02 +0800 Subject: [PATCH 2/9] update Advanced Experiment Tracking --- README.md | 2 +- README.zh-CN.md | 2 +- apps/backend/app/api/v1/endpoints/agent.py | 111 +++++- apps/backend/app/models/__init__.py | 2 + apps/backend/app/models/agent/__init__.py | 3 +- .../app/models/agent/optimization_jobs.py | 42 +++ .../app/repositories/agent/__init__.py | 4 +- .../agent/optimization_job_repository.py | 104 +++++- apps/backend/app/schemas/__init__.py | 4 + apps/backend/app/schemas/agent.py | 24 ++ .../app/services/agent/history_service.py | 322 +++++++++++++++++- .../tests/api/v1/endpoints/test_agent.py | 103 +++++- apps/frontend/src/api/agent.ts | 81 ++++- .../agent/components/AgentResultsCompare.tsx | 267 ++++++++++++++- .../src/features/agent/useAgentController.ts | 9 +- apps/frontend/src/pages/types.ts | 4 + 16 files changed, 1057 insertions(+), 27 deletions(-) diff --git a/README.md b/README.md index 04784de..d6d94d2 100644 --- a/README.md +++ b/README.md @@ -210,7 +210,7 @@ PRs welcome! Figaro is meant to be a readable, research-friendly FL platform. - [x] **Interactive Agent Planning** — Multi-turn dialogue support for refining experiments, plus visual topology previews (Plan Preview) before execution. - [x] **Execution Transparency** — Real-time tracking of node-level status during execution and automated natural-language interpretation of results. - [x] **Strict Configuration Engine** — Implement strict Pydantic/JSON Schema validation to resolve historical inconsistencies between `config_schema` and underlying algorithms. -- [ ] **Advanced Experiment Tracking** — Multi-dimensional search filtering (by metrics, hyperparameters, status) and configuration version control (diffing). +- [x] **Advanced Experiment Tracking** — Multi-dimensional search filtering (by metrics, hyperparameters, status) and configuration version control (diffing). **Phase 2: LLM & LoRA Federated Fine-Tuning** - [ ] **Native LLM Ecosystem Integration** — Seamless Hugging Face model loading (e.g., Llama 3, Qwen) and efficient parsing of JSONL instruction-tuning datasets. diff --git a/README.zh-CN.md b/README.zh-CN.md index 67c527d..a365f79 100644 --- a/README.zh-CN.md +++ b/README.zh-CN.md @@ -211,7 +211,7 @@ figaro/ - [x] **Agent 交互体验升级** —— 支持多轮对话微调实验计划,提供实验执行前的Plan Preview。 - [x] **执行与分析透明化** —— 支持实验节点的实时状态追踪,以及 Agent 驱动的运行结果自动化图表解释。 - [x] **配置引擎重构** —— 引入基于 Pydantic/JSON Schema 的严格强校验,彻底修复 `config_schema` 与底层算法实现不一致的问题。 -- [ ] **高阶实验管理** —— 支持按指标、超参等多维度搜索过滤实验历史,支持配置文件版本控制与 Diff 差异对比。 +- [x] **高阶实验管理** —— 支持按指标、超参等多维度搜索过滤实验历史,支持配置文件版本控制与 Diff 差异对比。 **Phase 2:LLM / LoRA 联邦微调支持** - [ ] **大模型生态原生接入** —— 内置 Hugging Face 适配层,一键加载主流开源模型,支持 JSONL 格式的指令微调数据集高效解析。 diff --git a/apps/backend/app/api/v1/endpoints/agent.py b/apps/backend/app/api/v1/endpoints/agent.py index db3130f..045c67b 100644 --- a/apps/backend/app/api/v1/endpoints/agent.py +++ b/apps/backend/app/api/v1/endpoints/agent.py @@ -4,12 +4,16 @@ from __future__ import annotations -from fastapi import APIRouter, status, HTTPException +from datetime import datetime + +from fastapi import APIRouter, status, HTTPException, Query import json from app.api.deps import AsyncSessionDep from app.schemas.agent import ( + AgentConfigDiffResponse, AgentConfigChangeResponse, + AgentConfigVersionResponse, AgentCurrentExperimentResponse, AgentCurrentPlanResponse, AgentExperimentResponse, @@ -247,12 +251,52 @@ def _build_history_summary(item) -> AgentOptimizationJobSummaryResponse: status_code=status.HTTP_200_OK, summary="List persisted agent optimization jobs", ) -async def list_optimization_jobs(session: AsyncSessionDep) -> list[AgentOptimizationJobSummaryResponse]: +async def list_optimization_jobs( + session: AsyncSessionDep, + job_status: str | None = Query(default=None, alias="status"), + q: str | None = Query(default=None), + model_name: str | None = Query(default=None), + objective: str | None = Query(default=None), + best_score_min: float | None = Query(default=None), + best_score_max: float | None = Query(default=None), + created_from: datetime | None = Query(default=None), + created_to: datetime | None = Query(default=None), + dataset: str | None = Query(default=None), + config_model: str | None = Query(default=None), + aggregation: str | None = Query(default=None), + num_clients: int | None = Query(default=None), + num_rounds: int | None = Query(default=None), + dataset_name_alias: str | None = Query(default=None, alias="dataset.name"), + config_model_alias: str | None = Query(default=None, alias="model.name"), + aggregation_alias: str | None = Query(default=None, alias="federated.aggregation"), + num_clients_alias: int | None = Query(default=None, alias="federated.num_clients"), + num_rounds_alias: int | None = Query(default=None, alias="federated.num_rounds"), +) -> list[AgentOptimizationJobSummaryResponse]: """ Return persisted Agent optimization jobs ordered by latest update time. """ service = AgentOptimizationHistoryService(session) - jobs = await service.list_jobs() + config_filters = { + "dataset.name": dataset_name_alias or dataset, + "model.name": config_model_alias or config_model, + "federated.aggregation": aggregation_alias or aggregation, + "federated.num_clients": num_clients_alias if num_clients_alias is not None else num_clients, + "federated.num_rounds": num_rounds_alias if num_rounds_alias is not None else num_rounds, + } + try: + jobs = await service.list_jobs( + status=job_status, + q=q, + model_name=model_name, + objective=objective, + best_score_min=best_score_min, + best_score_max=best_score_max, + created_from=created_from, + created_to=created_to, + config_filters=config_filters, + ) + except ValueError as exc: + raise HTTPException(status_code=422, detail=str(exc)) from exc return [_build_history_summary(item) for item in jobs] @@ -290,6 +334,67 @@ async def get_optimization_job( return _build_progress_response(snapshot) +def _to_config_version_response(version) -> AgentConfigVersionResponse: + return AgentConfigVersionResponse( + id=version.id, + optimization_job_id=version.optimization_job_id, + run_id=version.run_id, + iteration=version.iteration, + source=version.source, + label=version.label, + config_hash=version.config_hash, + config_json=version.config_json if isinstance(version.config_json, dict) else {}, + diff_json=[ + AgentConfigChangeResponse.model_validate(item) + for item in (version.diff_json if isinstance(version.diff_json, list) else []) + ], + created_at=version.created_at, + ) + + +@agent_router.get( + "/optimization-jobs/{optimization_job_id}/config-versions", + response_model=list[AgentConfigVersionResponse], + status_code=status.HTTP_200_OK, + summary="List persisted config versions for an agent optimization job", +) +async def list_optimization_job_config_versions( + optimization_job_id: int, + session: AsyncSessionDep, +) -> list[AgentConfigVersionResponse]: + """Return config versions captured during planning and experiment execution.""" + service = AgentOptimizationHistoryService(session) + versions = await service.list_config_versions(optimization_job_id) + return [_to_config_version_response(version) for version in versions] + + +@agent_router.get( + "/optimization-jobs/{optimization_job_id}/config-diff", + response_model=AgentConfigDiffResponse, + status_code=status.HTTP_200_OK, + summary="Diff two persisted config versions for an agent optimization job", +) +async def diff_optimization_job_config_versions( + optimization_job_id: int, + session: AsyncSessionDep, + to_version_id: int = Query(...), + from_version_id: int | None = Query(default=None), +) -> AgentConfigDiffResponse: + """Return a normalized config diff between two persisted versions.""" + service = AgentOptimizationHistoryService(session) + changes = await service.diff_config_versions( + optimization_job_id=optimization_job_id, + from_version_id=from_version_id, + to_version_id=to_version_id, + ) + return AgentConfigDiffResponse( + optimization_job_id=optimization_job_id, + from_version_id=from_version_id, + to_version_id=to_version_id, + changes=[AgentConfigChangeResponse.model_validate(item) for item in changes], + ) + + @agent_router.post( "/optimize/start", response_model=AgentOptimizeProgressResponse, diff --git a/apps/backend/app/models/__init__.py b/apps/backend/app/models/__init__.py index e8852bb..9155124 100644 --- a/apps/backend/app/models/__init__.py +++ b/apps/backend/app/models/__init__.py @@ -1,8 +1,10 @@ """Model package marker — import all models so SQLModel registers them.""" from app.models.agent import ( # noqa: F401 + AgentConfigVersion, AgentExperiment, AgentExperimentRun, AgentExperimentRunLog, AgentExperimentRunResult, + AgentOptimizationJob, ) diff --git a/apps/backend/app/models/agent/__init__.py b/apps/backend/app/models/agent/__init__.py index 3f6f494..8fe0be4 100644 --- a/apps/backend/app/models/agent/__init__.py +++ b/apps/backend/app/models/agent/__init__.py @@ -10,7 +10,7 @@ AgentExperimentRunStatus, AgentExperimentStatus, ) -from app.models.agent.optimization_jobs import AgentOptimizationJob, AgentOptimizationJobStatus +from app.models.agent.optimization_jobs import AgentConfigVersion, AgentOptimizationJob, AgentOptimizationJobStatus __all__ = [ "AgentExperiment", @@ -21,4 +21,5 @@ "AgentExperimentRunResult", "AgentOptimizationJob", "AgentOptimizationJobStatus", + "AgentConfigVersion", ] diff --git a/apps/backend/app/models/agent/optimization_jobs.py b/apps/backend/app/models/agent/optimization_jobs.py index 8848516..650c05b 100644 --- a/apps/backend/app/models/agent/optimization_jobs.py +++ b/apps/backend/app/models/agent/optimization_jobs.py @@ -93,3 +93,45 @@ class AgentOptimizationJob(SQLModel, table=True): default=None, sa_column=Column(DateTime(timezone=True), nullable=True), ) + + +class AgentConfigVersion(SQLModel, table=True): + __tablename__ = "agent_config_versions" + + id: int | None = Field( + default=None, + sa_column=Column(Integer, primary_key=True, nullable=False), + ) + optimization_job_id: int = Field( + sa_column=Column(Integer, nullable=False, index=True), + ) + run_id: str | None = Field( + default=None, + sa_column=Column(String(64), nullable=True, index=True), + ) + iteration: int = Field( + default=0, + sa_column=Column(Integer, nullable=False, index=True), + ) + source: str = Field( + default="experiment", + sa_column=Column(String(32), nullable=False, index=True), + ) + label: str = Field( + sa_column=Column(String(255), nullable=False), + ) + config_hash: str = Field( + sa_column=Column(String(64), nullable=False, index=True), + ) + config_json: dict[str, Any] = Field( + default_factory=dict, + sa_column=Column(JSON().with_variant(JSONB, "postgresql"), nullable=False), + ) + diff_json: list[dict[str, Any]] = Field( + default_factory=list, + sa_column=Column(JSON().with_variant(JSONB, "postgresql"), nullable=False), + ) + created_at: datetime = Field( + default_factory=utcnow, + sa_column=Column(DateTime(timezone=True), nullable=False), + ) diff --git a/apps/backend/app/repositories/agent/__init__.py b/apps/backend/app/repositories/agent/__init__.py index dec957e..c02fe11 100644 --- a/apps/backend/app/repositories/agent/__init__.py +++ b/apps/backend/app/repositories/agent/__init__.py @@ -1,6 +1,6 @@ """Agent repository exports.""" from app.repositories.agent.experiment_repository import AgentExperimentRepository -from app.repositories.agent.optimization_job_repository import AgentOptimizationJobRepository +from app.repositories.agent.optimization_job_repository import AgentConfigVersionRepository, AgentOptimizationJobRepository -__all__ = ["AgentExperimentRepository", "AgentOptimizationJobRepository"] +__all__ = ["AgentExperimentRepository", "AgentOptimizationJobRepository", "AgentConfigVersionRepository"] diff --git a/apps/backend/app/repositories/agent/optimization_job_repository.py b/apps/backend/app/repositories/agent/optimization_job_repository.py index 6fa8f95..b014580 100644 --- a/apps/backend/app/repositories/agent/optimization_job_repository.py +++ b/apps/backend/app/repositories/agent/optimization_job_repository.py @@ -4,13 +4,14 @@ from __future__ import annotations +from datetime import datetime from typing import Any -from sqlalchemy import func +from sqlalchemy import func, or_ from sqlalchemy.ext.asyncio import AsyncSession from sqlmodel import select -from app.models.agent import AgentOptimizationJob, AgentOptimizationJobStatus +from app.models.agent import AgentConfigVersion, AgentOptimizationJob, AgentOptimizationJobStatus from app.models.base import utcnow @@ -62,8 +63,38 @@ async def get_job_by_name(self, job_name: str) -> AgentOptimizationJob | None: result = await self.session.execute(stmt) return result.scalars().first() - async def list_jobs(self) -> list[AgentOptimizationJob]: + async def list_jobs( + self, + *, + status: AgentOptimizationJobStatus | None = None, + q: str | None = None, + model_name: str | None = None, + best_score_min: float | None = None, + best_score_max: float | None = None, + created_from: datetime | None = None, + created_to: datetime | None = None, + ) -> list[AgentOptimizationJob]: stmt = select(AgentOptimizationJob).order_by(AgentOptimizationJob.updated_at.desc(), AgentOptimizationJob.id.desc()) + if status is not None: + stmt = stmt.where(AgentOptimizationJob.status == status) + if q: + like_value = f"%{q.lower()}%" + stmt = stmt.where( + or_( + func.lower(AgentOptimizationJob.goal).like(like_value), + func.lower(AgentOptimizationJob.job_name).like(like_value), + ) + ) + if model_name: + stmt = stmt.where(func.lower(AgentOptimizationJob.model_name) == model_name.lower()) + if best_score_min is not None: + stmt = stmt.where(AgentOptimizationJob.best_score >= best_score_min) + if best_score_max is not None: + stmt = stmt.where(AgentOptimizationJob.best_score <= best_score_max) + if created_from is not None: + stmt = stmt.where(AgentOptimizationJob.created_at >= created_from) + if created_to is not None: + stmt = stmt.where(AgentOptimizationJob.created_at <= created_to) result = await self.session.execute(stmt) return list(result.scalars().all()) @@ -108,3 +139,70 @@ async def update_job( self.session.add(job) await self.session.flush() return job + + +class AgentConfigVersionRepository: + """CRUD helpers for persisted agent configuration versions.""" + + def __init__(self, session: AsyncSession): + self.session = session + + async def create_version( + self, + *, + optimization_job_id: int, + run_id: str | None, + iteration: int, + source: str, + label: str, + config_hash: str, + config_json: dict[str, Any], + diff_json: list[dict[str, Any]], + ) -> AgentConfigVersion: + version = AgentConfigVersion( + optimization_job_id=optimization_job_id, + run_id=run_id, + iteration=iteration, + source=source, + label=label, + config_hash=config_hash, + config_json=config_json, + diff_json=diff_json, + ) + self.session.add(version) + await self.session.flush() + return version + + async def find_version( + self, + *, + optimization_job_id: int, + run_id: str | None, + iteration: int, + source: str, + config_hash: str, + ) -> AgentConfigVersion | None: + stmt = select(AgentConfigVersion).where( + AgentConfigVersion.optimization_job_id == optimization_job_id, + AgentConfigVersion.iteration == iteration, + AgentConfigVersion.source == source, + AgentConfigVersion.config_hash == config_hash, + ) + if run_id is None: + stmt = stmt.where(AgentConfigVersion.run_id.is_(None)) + else: + stmt = stmt.where(AgentConfigVersion.run_id == run_id) + result = await self.session.execute(stmt) + return result.scalars().first() + + async def get_version(self, version_id: int) -> AgentConfigVersion | None: + return await self.session.get(AgentConfigVersion, version_id) + + async def list_versions(self, optimization_job_id: int) -> list[AgentConfigVersion]: + stmt = ( + select(AgentConfigVersion) + .where(AgentConfigVersion.optimization_job_id == optimization_job_id) + .order_by(AgentConfigVersion.created_at.asc(), AgentConfigVersion.id.asc()) + ) + result = await self.session.execute(stmt) + return list(result.scalars().all()) diff --git a/apps/backend/app/schemas/__init__.py b/apps/backend/app/schemas/__init__.py index b793e9f..ad3ca51 100644 --- a/apps/backend/app/schemas/__init__.py +++ b/apps/backend/app/schemas/__init__.py @@ -29,6 +29,8 @@ ) from app.schemas.message import Message from app.schemas.agent import ( + AgentConfigDiffResponse, + AgentConfigVersionResponse, AgentCurrentPlanResponse, AgentExperimentSummary, AgentOptimizationJobSummaryResponse, @@ -66,4 +68,6 @@ "AgentOptimizationJobSummaryResponse", "AgentCurrentPlanResponse", "AgentExperimentSummary", + "AgentConfigVersionResponse", + "AgentConfigDiffResponse", ] diff --git a/apps/backend/app/schemas/agent.py b/apps/backend/app/schemas/agent.py index 3baec67..2425c04 100644 --- a/apps/backend/app/schemas/agent.py +++ b/apps/backend/app/schemas/agent.py @@ -67,6 +67,30 @@ class AgentConfigChangeResponse(BaseModel): new_value: Any = None +class AgentConfigVersionResponse(BaseModel): + """One persisted config version for a historical Agent optimization.""" + + id: int + optimization_job_id: int + run_id: str | None = None + iteration: int + source: str + label: str + config_hash: str + config_json: dict[str, Any] = Field(default_factory=dict) + diff_json: list[AgentConfigChangeResponse] = Field(default_factory=list) + created_at: datetime + + +class AgentConfigDiffResponse(BaseModel): + """Diff payload between two persisted config versions.""" + + optimization_job_id: int + from_version_id: int | None = None + to_version_id: int + changes: list[AgentConfigChangeResponse] = Field(default_factory=list) + + class AgentCurrentPlanResponse(BaseModel): """Planner output for the iteration currently being prepared or executed.""" diff --git a/apps/backend/app/services/agent/history_service.py b/apps/backend/app/services/agent/history_service.py index 7dce611..686a1a2 100644 --- a/apps/backend/app/services/agent/history_service.py +++ b/apps/backend/app/services/agent/history_service.py @@ -5,6 +5,8 @@ from __future__ import annotations import copy +import hashlib +import json from datetime import datetime from typing import Any @@ -12,9 +14,10 @@ from sqlalchemy.ext.asyncio import AsyncSession from app.core import exceptions -from app.models.agent import AgentOptimizationJob, AgentOptimizationJobStatus -from app.repositories.agent import AgentOptimizationJobRepository +from app.models.agent import AgentConfigVersion, AgentOptimizationJob, AgentOptimizationJobStatus +from app.repositories.agent import AgentConfigVersionRepository, AgentOptimizationJobRepository +from .memory import compute_config_diff from .summary import get_last_global_accuracy @@ -26,6 +29,7 @@ class AgentOptimizationHistoryService: def __init__(self, session: AsyncSession): self.session = session self.repository = AgentOptimizationJobRepository(session) + self.config_version_repository = AgentConfigVersionRepository(session) async def create_job( self, @@ -47,6 +51,7 @@ async def create_job( resolved_status = AgentOptimizationJobStatus(status.lower()) try: + safe_snapshot = self._to_json_safe(snapshot) job = await self.repository.create_job( task_id=task_id, job_name=job_name, @@ -55,8 +60,9 @@ async def create_job( model_name=model_name, status=resolved_status, max_iterations=max_iterations, - snapshot_json=self._to_json_safe(snapshot), + snapshot_json=safe_snapshot, ) + await self._sync_config_versions(job, safe_snapshot) await self.session.commit() await self.session.refresh(job) return job @@ -64,8 +70,38 @@ async def create_job( await self.session.rollback() raise exceptions.JobAlreadyExists(self.JOB_NAME_CONFLICT_MESSAGE) from exc - async def list_jobs(self) -> list[AgentOptimizationJob]: - return await self.repository.list_jobs() + async def list_jobs( + self, + *, + status: AgentOptimizationJobStatus | str | None = None, + q: str | None = None, + model_name: str | None = None, + objective: str | None = None, + best_score_min: float | None = None, + best_score_max: float | None = None, + created_from: datetime | None = None, + created_to: datetime | None = None, + config_filters: dict[str, Any] | None = None, + ) -> list[AgentOptimizationJob]: + resolved_status = self._resolve_status(status) + jobs = await self.repository.list_jobs( + status=resolved_status, + q=q, + model_name=model_name, + best_score_min=best_score_min, + best_score_max=best_score_max, + created_from=created_from, + created_to=created_to, + ) + return [ + job + for job in jobs + if self._matches_snapshot_filters( + job.snapshot_json if isinstance(job.snapshot_json, dict) else {}, + objective=objective, + config_filters=config_filters or {}, + ) + ] async def get_job_or_raise(self, optimization_job_id: int) -> AgentOptimizationJob: job = await self.repository.get_job(optimization_job_id) @@ -96,6 +132,7 @@ async def update_job_snapshot( best_score = get_last_global_accuracy(best_metrics) if isinstance(best_metrics, dict) else None try: + safe_snapshot = self._to_json_safe(snapshot) await self.repository.update_job( job, job_name=str(job_name) if job_name is not None else None, @@ -105,9 +142,10 @@ async def update_job_snapshot( completed_iterations=int(snapshot.get("completed_iterations") or 0), simulation_job_id=self._extract_simulation_job_id(snapshot), best_score=best_score, - snapshot_json=self._to_json_safe(snapshot), + snapshot_json=safe_snapshot, finished_at=finished_at, ) + await self._sync_config_versions(job, safe_snapshot) await self.session.commit() await self.session.refresh(job) return job @@ -132,6 +170,278 @@ async def _ensure_unique_job_name(self, job_name: str, *, exclude_job_id: int | await self.session.flush() else: raise exceptions.JobAlreadyExists(self.JOB_NAME_CONFLICT_MESSAGE) + + async def list_config_versions(self, optimization_job_id: int) -> list[AgentConfigVersion]: + job = await self.get_job_or_raise(optimization_job_id) + versions = await self.config_version_repository.list_versions(optimization_job_id) + if versions: + return versions + snapshot = job.snapshot_json if isinstance(job.snapshot_json, dict) else {} + if not self._build_config_version_candidates(snapshot): + return [] + await self._sync_config_versions(job, snapshot) + await self.session.commit() + return await self.config_version_repository.list_versions(optimization_job_id) + + async def get_config_version_or_raise(self, optimization_job_id: int, version_id: int) -> AgentConfigVersion: + await self.get_job_or_raise(optimization_job_id) + version = await self.config_version_repository.get_version(version_id) + if version is None or version.optimization_job_id != optimization_job_id: + raise exceptions.ResourceNotFound("Agent config version not found") + return version + + async def diff_config_versions( + self, + *, + optimization_job_id: int, + from_version_id: int | None, + to_version_id: int, + ) -> list[dict[str, Any]]: + to_version = await self.get_config_version_or_raise(optimization_job_id, to_version_id) + previous_config = None + if from_version_id is not None: + from_version = await self.get_config_version_or_raise(optimization_job_id, from_version_id) + previous_config = from_version.config_json if isinstance(from_version.config_json, dict) else {} + changes = compute_config_diff(previous_config, to_version.config_json if isinstance(to_version.config_json, dict) else {}) + return [self._serialize_config_change(change) for change in changes] + + async def _sync_config_versions(self, job: AgentOptimizationJob, snapshot: dict[str, Any]) -> None: + if job.id is None: + return + for candidate in self._build_config_version_candidates(snapshot): + existing = await self.config_version_repository.find_version( + optimization_job_id=job.id, + run_id=candidate["run_id"], + iteration=candidate["iteration"], + source=candidate["source"], + config_hash=candidate["config_hash"], + ) + if existing is not None: + continue + await self.config_version_repository.create_version( + optimization_job_id=job.id, + run_id=candidate["run_id"], + iteration=candidate["iteration"], + source=candidate["source"], + label=candidate["label"], + config_hash=candidate["config_hash"], + config_json=candidate["config_json"], + diff_json=candidate["diff_json"], + ) + + @classmethod + def _build_config_version_candidates(cls, snapshot: dict[str, Any]) -> list[dict[str, Any]]: + candidates: list[dict[str, Any]] = [] + previous_config: dict[str, Any] | None = None + + draft_experiments = snapshot.get("draft_experiments") + if isinstance(draft_experiments, list): + for index, item in enumerate(draft_experiments, start=1): + if not isinstance(item, dict): + continue + config = item.get("config_patch") or item.get("config") + if not isinstance(config, dict): + continue + candidates.append( + cls._build_config_version_candidate( + config=config, + previous_config=previous_config, + iteration=int(item.get("iteration") or index), + source="draft", + run_id=None, + label=str(item.get("name") or f"Draft {index}"), + diff_json=item.get("config_diff"), + ) + ) + previous_config = config + + experiments = snapshot.get("experiments") + if isinstance(experiments, list): + previous_config = None + for index, item in enumerate(experiments, start=1): + if not isinstance(item, dict): + continue + config = item.get("config") + if not isinstance(config, dict): + continue + candidates.append( + cls._build_config_version_candidate( + config=config, + previous_config=previous_config, + iteration=int(item.get("iteration") or index), + source="experiment", + run_id=str(item.get("run_id")) if item.get("run_id") is not None else None, + label=str(item.get("name") or item.get("plan_summary") or f"Experiment {index}"), + diff_json=item.get("config_diff"), + ) + ) + previous_config = config + + best_config = snapshot.get("best_config") + if isinstance(best_config, dict): + best_record = cls._best_snapshot_experiment(snapshot) + candidates.append( + cls._build_config_version_candidate( + config=best_config, + previous_config=previous_config, + iteration=int(best_record.get("iteration") or 0) if best_record else 0, + source="best", + run_id=str(best_record.get("run_id")) if best_record and best_record.get("run_id") is not None else None, + label="Best configuration", + diff_json=None, + ) + ) + + return candidates + + @classmethod + def _build_config_version_candidate( + cls, + *, + config: dict[str, Any], + previous_config: dict[str, Any] | None, + iteration: int, + source: str, + run_id: str | None, + label: str, + diff_json: Any, + ) -> dict[str, Any]: + safe_config = cls._to_json_safe(config) + if isinstance(diff_json, list): + safe_diff = cls._to_json_safe(diff_json) + else: + safe_diff = [ + cls._serialize_config_change(change) + for change in compute_config_diff(previous_config, safe_config) + ] + return { + "run_id": run_id, + "iteration": iteration, + "source": source, + "label": label[:255], + "config_hash": cls._config_hash(safe_config), + "config_json": safe_config, + "diff_json": safe_diff, + } + + @staticmethod + def _best_snapshot_experiment(snapshot: dict[str, Any]) -> dict[str, Any] | None: + experiments = snapshot.get("experiments") + if not isinstance(experiments, list): + return None + best_record = None + best_score = None + for item in experiments: + if not isinstance(item, dict): + continue + score = item.get("score") + if score is None: + continue + try: + numeric_score = float(score) + except (TypeError, ValueError): + continue + if best_score is None or numeric_score > best_score: + best_record = item + best_score = numeric_score + return best_record + + @staticmethod + def _serialize_config_change(change: Any) -> dict[str, Any]: + if isinstance(change, dict): + return { + "path": str(change.get("path", "$")), + "change_type": str(change.get("change_type", "updated")), + "old_value": copy.deepcopy(change.get("old_value")), + "new_value": copy.deepcopy(change.get("new_value")), + } + return { + "path": change.path, + "change_type": change.change_type, + "old_value": copy.deepcopy(change.old_value), + "new_value": copy.deepcopy(change.new_value), + } + + @staticmethod + def _config_hash(config: dict[str, Any]) -> str: + payload = json.dumps(config, sort_keys=True, separators=(",", ":"), default=str) + return hashlib.sha256(payload.encode("utf-8")).hexdigest() + + @classmethod + def _resolve_status(cls, status: AgentOptimizationJobStatus | str | None) -> AgentOptimizationJobStatus | None: + if status is None: + return None + if isinstance(status, AgentOptimizationJobStatus): + return status + status_text = str(status).strip().lower() + if not status_text or status_text == "all": + return None + return AgentOptimizationJobStatus(status_text) + + @classmethod + def _matches_snapshot_filters( + cls, + snapshot: dict[str, Any], + *, + objective: str | None, + config_filters: dict[str, Any], + ) -> bool: + if objective and str(snapshot.get("objective", "")).lower() != objective.lower(): + return False + for path, expected in config_filters.items(): + if expected is None or str(expected).strip() == "": + continue + if not cls._snapshot_has_config_value(snapshot, path, expected): + return False + return True + + @classmethod + def _snapshot_has_config_value(cls, snapshot: dict[str, Any], path: str, expected: Any) -> bool: + for config in cls._iter_snapshot_configs(snapshot): + value = cls._get_value_by_path(config, path) + if value is not None and cls._values_equal(value, expected): + return True + return False + + @classmethod + def _iter_snapshot_configs(cls, snapshot: dict[str, Any]): + for key in ("best_config", "config_constraints"): + value = snapshot.get(key) + if isinstance(value, dict): + yield value + current_experiment = snapshot.get("current_experiment") + if isinstance(current_experiment, dict) and isinstance(current_experiment.get("config"), dict): + yield current_experiment["config"] + for collection_key in ("experiments", "draft_experiments"): + collection = snapshot.get(collection_key) + if not isinstance(collection, list): + continue + for item in collection: + if not isinstance(item, dict): + continue + for config_key in ("config", "config_patch"): + config = item.get(config_key) + if isinstance(config, dict): + yield config + + @staticmethod + def _get_value_by_path(obj: dict[str, Any], path: str) -> Any: + current: Any = obj + for part in path.split("."): + if not isinstance(current, dict) or part not in current: + return None + current = current[part] + return current + + @staticmethod + def _values_equal(actual: Any, expected: Any) -> bool: + if isinstance(actual, (int, float)) and isinstance(expected, str): + try: + return float(actual) == float(expected) + except ValueError: + return False + return str(actual).lower() == str(expected).lower() + @staticmethod def _extract_simulation_job_id(snapshot: dict[str, Any]) -> int | None: current_experiment = snapshot.get("current_experiment") diff --git a/apps/backend/tests/api/v1/endpoints/test_agent.py b/apps/backend/tests/api/v1/endpoints/test_agent.py index cf62080..dbf50cb 100644 --- a/apps/backend/tests/api/v1/endpoints/test_agent.py +++ b/apps/backend/tests/api/v1/endpoints/test_agent.py @@ -370,7 +370,7 @@ class _FakeHistoryService: def __init__(self, _session): self._session = _session - async def list_jobs(self): + async def list_jobs(self, **_filters): return [history_job] async def get_job_or_raise(self, optimization_job_id): @@ -393,3 +393,104 @@ async def get_job_or_raise(self, optimization_job_id): assert detail_payload["optimization_job_id"] == 7 assert detail_payload["job_name"] == "cifar10-alpha-sweep" assert detail_payload["resolved_objective"] == "accuracy" + + +def test_agent_optimization_jobs_history_filters_are_forwarded(client, monkeypatch): + import app.api.v1.endpoints.agent as agent_module + + captured = {} + + class _FakeHistoryService: + def __init__(self, _session): + self._session = _session + + async def list_jobs(self, **filters): + captured.update(filters) + return [] + + monkeypatch.setattr(agent_module, "AgentOptimizationHistoryService", _FakeHistoryService) + + response = client.get( + "/api/v1/agent/optimization-jobs" + "?status=completed" + "&q=cifar" + "&model_name=gpt-test" + "&objective=accuracy" + "&best_score_min=0.8" + "&best_score_max=0.95" + "&dataset=CIFAR-10" + "&config_model=FedAvgCNN" + "&aggregation=fedavg" + "&num_clients=3" + "&num_rounds=10" + ) + + assert response.status_code == 200 + assert captured["status"] == "completed" + assert captured["q"] == "cifar" + assert captured["model_name"] == "gpt-test" + assert captured["objective"] == "accuracy" + assert captured["best_score_min"] == 0.8 + assert captured["best_score_max"] == 0.95 + assert captured["config_filters"] == { + "dataset.name": "CIFAR-10", + "model.name": "FedAvgCNN", + "federated.aggregation": "fedavg", + "federated.num_clients": 3, + "federated.num_rounds": 10, + } + + +def test_agent_config_version_endpoints(client, monkeypatch): + import app.api.v1.endpoints.agent as agent_module + + now = datetime.now(timezone.utc) + version = SimpleNamespace( + id=3, + optimization_job_id=7, + run_id="run-1", + iteration=1, + source="experiment", + label="alpha 0.1", + config_hash="abc123", + config_json={"dataset": {"alpha": 0.1}}, + diff_json=[ + { + "path": "dataset.alpha", + "change_type": "updated", + "old_value": 0.5, + "new_value": 0.1, + } + ], + created_at=now, + ) + + class _FakeHistoryService: + def __init__(self, _session): + self._session = _session + + async def list_config_versions(self, optimization_job_id): + assert optimization_job_id == 7 + return [version] + + async def diff_config_versions(self, *, optimization_job_id, from_version_id, to_version_id): + assert optimization_job_id == 7 + assert from_version_id is None + assert to_version_id == 3 + return version.diff_json + + monkeypatch.setattr(agent_module, "AgentOptimizationHistoryService", _FakeHistoryService) + + versions_response = client.get("/api/v1/agent/optimization-jobs/7/config-versions") + assert versions_response.status_code == 200 + versions_payload = versions_response.json() + assert versions_payload[0]["id"] == 3 + assert versions_payload[0]["label"] == "alpha 0.1" + assert versions_payload[0]["diff_json"][0]["path"] == "dataset.alpha" + + diff_response = client.get("/api/v1/agent/optimization-jobs/7/config-diff?to_version_id=3") + assert diff_response.status_code == 200 + diff_payload = diff_response.json() + assert diff_payload["optimization_job_id"] == 7 + assert diff_payload["to_version_id"] == 3 + assert diff_payload["changes"][0]["new_value"] == 0.1 diff --git a/apps/frontend/src/api/agent.ts b/apps/frontend/src/api/agent.ts index 0e0cda6..876d06d 100644 --- a/apps/frontend/src/api/agent.ts +++ b/apps/frontend/src/api/agent.ts @@ -21,6 +21,42 @@ export type AgentConfigChange = { new_value: unknown; }; +export type AgentHistoryFilters = { + q?: string; + status?: string; + model_name?: string; + objective?: AgentOptimizationObjective | "all"; + best_score_min?: string; + best_score_max?: string; + created_from?: string; + created_to?: string; + dataset?: string; + config_model?: string; + aggregation?: string; + num_clients?: string; + num_rounds?: string; +}; + +export type AgentConfigVersion = { + id: number; + optimization_job_id: number; + run_id: string | null; + iteration: number; + source: string; + label: string; + config_hash: string; + config_json: Record; + diff_json: AgentConfigChange[]; + created_at: string; +}; + +export type AgentConfigDiffResponse = { + optimization_job_id: number; + from_version_id: number | null; + to_version_id: number; + changes: AgentConfigChange[]; +}; + export type AgentCurrentPlan = { iteration: number; iteration_goal: string; @@ -192,6 +228,33 @@ async function readJson(response: Response): Promise { return (await response.json()) as T; } +function addOptionalParam(params: URLSearchParams, key: string, value: unknown): void { + if (value === undefined || value === null) return; + const text = String(value).trim(); + if (!text || text === "all") return; + params.set(key, text); +} + +function historyFiltersToParams(filters?: AgentHistoryFilters): string { + const params = new URLSearchParams(); + if (!filters) return ""; + addOptionalParam(params, "q", filters.q); + addOptionalParam(params, "status", filters.status); + addOptionalParam(params, "model_name", filters.model_name); + addOptionalParam(params, "objective", filters.objective); + addOptionalParam(params, "best_score_min", filters.best_score_min); + addOptionalParam(params, "best_score_max", filters.best_score_max); + addOptionalParam(params, "created_from", filters.created_from); + addOptionalParam(params, "created_to", filters.created_to); + addOptionalParam(params, "dataset", filters.dataset); + addOptionalParam(params, "config_model", filters.config_model); + addOptionalParam(params, "aggregation", filters.aggregation); + addOptionalParam(params, "num_clients", filters.num_clients); + addOptionalParam(params, "num_rounds", filters.num_rounds); + const query = params.toString(); + return query ? `?${query}` : ""; +} + export const agentApi = { async getConfigSchema(): Promise> { const response = await fetch(`${baseUrl}/api/v1/agent/config/schema`); @@ -226,8 +289,8 @@ export const agentApi = { return readJson(response); }, - async listOptimizationJobs(): Promise { - const response = await fetch(`${baseUrl}/api/v1/agent/optimization-jobs`); + async listOptimizationJobs(filters?: AgentHistoryFilters): Promise { + const response = await fetch(`${baseUrl}/api/v1/agent/optimization-jobs${historyFiltersToParams(filters)}`); return readJson(response); }, @@ -303,4 +366,18 @@ export const agentApi = { } return response.json(); }, + + async listConfigVersions(optimizationJobId: number): Promise { + const response = await fetch(`${baseUrl}/api/v1/agent/optimization-jobs/${encodeURIComponent(optimizationJobId)}/config-versions`); + return readJson(response); + }, + + async getConfigDiff(optimizationJobId: number, toVersionId: number, fromVersionId?: number | null): Promise { + const params = new URLSearchParams({ to_version_id: String(toVersionId) }); + if (fromVersionId != null) { + params.set("from_version_id", String(fromVersionId)); + } + const response = await fetch(`${baseUrl}/api/v1/agent/optimization-jobs/${encodeURIComponent(optimizationJobId)}/config-diff?${params.toString()}`); + return readJson(response); + }, }; diff --git a/apps/frontend/src/features/agent/components/AgentResultsCompare.tsx b/apps/frontend/src/features/agent/components/AgentResultsCompare.tsx index a9d608d..6be4ca7 100644 --- a/apps/frontend/src/features/agent/components/AgentResultsCompare.tsx +++ b/apps/frontend/src/features/agent/components/AgentResultsCompare.tsx @@ -1,14 +1,17 @@ import { useEffect, useMemo, useState } from "react"; -import { History, Bot, Activity, Trophy, ArrowRight, Target, Calendar, FileJson, Settings2 } from "lucide-react"; +import { History, Bot, Activity, Trophy, ArrowRight, Target, Calendar, FileJson, Settings2, Search, RotateCcw, GitCompare } from "lucide-react"; import { Card, CardContent, CardHeader, CardTitle, CardDescription } from "../../../components/ui/card"; import { Badge } from "../../../components/ui/badge"; import { Button } from "../../../components/ui/button"; +import { Input } from "../../../components/ui/input"; +import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "../../../components/ui/select"; import { Separator } from "../../../components/ui/separator"; import { MiniLineChart } from "../../simulation/components/MiniLineChart"; import { getValueByPath, isRecord } from "../../simulation/utils"; import { fmt } from "../../../lib/time"; import type { AgentPageProps } from "../../../pages/types"; -import { baseUrl } from "../../../../src/api/client"; +import { baseUrl } from "../../../api/client"; +import { agentApi, type AgentConfigChange, type AgentConfigVersion, type AgentHistoryFilters } from "../../../api/agent"; import ReactMarkdown from 'react-markdown'; import remarkGfm from 'remark-gfm'; import { collectAgentSchemaFields, formatFieldValue, optionLabel, type AgentSchemaField } from "../schema"; @@ -35,6 +38,21 @@ type ConfigSummaryItem = { value: string; }; +const HISTORY_STATUS_OPTIONS = [ + { value: "all", label: "All status" }, + { value: "completed", label: "Completed" }, + { value: "running", label: "Running" }, + { value: "pending_review", label: "Pending review" }, + { value: "failed", label: "Failed" }, + { value: "queued", label: "Queued" }, +]; + +const OBJECTIVE_OPTIONS = [ + { value: "all", label: "All objectives" }, + { value: "auto", label: "Auto" }, + { value: "accuracy", label: "Accuracy" }, +]; + function asConfigRecord(value: unknown): Record | null { return isRecord(value) ? value : null; } @@ -80,12 +98,38 @@ function buildConfigSummary( .filter((item): item is ConfigSummaryItem => item !== null); } +function formatDiffValue(value: unknown): string { + if (value === null || value === undefined) return "-"; + if (typeof value === "string") return value; + return JSON.stringify(value); +} + +function diffBadgeClass(changeType: string): string { + if (changeType === "added" || changeType === "initialize") return "border-emerald-500/70 text-emerald-600"; + if (changeType === "removed") return "border-red-500/70 text-red-600"; + return "border-amber-500/70 text-amber-600"; +} + export function AgentResultsCompare(props: AgentPageProps) { - const { historyJobs, selectedHistory, selectHistoryJob, setWorkflowStep, configSchema } = props; + const { + historyFilters, + historyJobs, + selectedHistory, + selectHistoryJob, + setHistoryFilters, + refreshHistoryJobs, + setWorkflowStep, + configSchema, + notifyError, + } = props; - const jobs = (historyJobs || []).filter((job: any) => job.status === "completed"); + const jobs = historyJobs || []; const [selectedRunId, setSelectedRunId] = useState(null); const [metrics, setMetrics] = useState(null); + const [configVersions, setConfigVersions] = useState([]); + const [fromVersionKey, setFromVersionKey] = useState("baseline"); + const [toVersionKey, setToVersionKey] = useState(""); + const [configDiff, setConfigDiff] = useState([]); const experiments = selectedHistory?.experiments || []; const bestExp = experiments.reduce((prev: any, curr: any) => @@ -122,6 +166,83 @@ export function AgentResultsCompare(props: AgentPageProps) { .catch(() => {}); }, [selectedRunId]); + useEffect(() => { + const optimizationJobId = selectedHistory?.optimization_job_id; + if (!optimizationJobId) { + setConfigVersions([]); + setFromVersionKey("baseline"); + setToVersionKey(""); + setConfigDiff([]); + return; + } + + let cancelled = false; + agentApi.listConfigVersions(optimizationJobId) + .then((versions) => { + if (cancelled) return; + setConfigVersions(versions); + setFromVersionKey(versions.length > 1 ? String(versions[0].id) : "baseline"); + setToVersionKey(versions.length > 0 ? String(versions[versions.length - 1].id) : ""); + }) + .catch((error) => { + if (!cancelled) { + setConfigVersions([]); + setConfigDiff([]); + notifyError(error, "agent-config-versions"); + } + }); + + return () => { + cancelled = true; + }; + }, [selectedHistory?.optimization_job_id]); + + useEffect(() => { + const optimizationJobId = selectedHistory?.optimization_job_id; + const toVersionId = Number(toVersionKey); + if (!optimizationJobId || !toVersionKey || !Number.isFinite(toVersionId)) { + setConfigDiff([]); + return; + } + + const fromVersionId = fromVersionKey === "baseline" ? null : Number(fromVersionKey); + let cancelled = false; + agentApi.getConfigDiff( + optimizationJobId, + toVersionId, + fromVersionId !== null && Number.isFinite(fromVersionId) && fromVersionId !== toVersionId ? fromVersionId : null, + ) + .then((payload) => { + if (!cancelled) { + setConfigDiff(payload.changes); + } + }) + .catch((error) => { + if (!cancelled) { + setConfigDiff([]); + notifyError(error, "agent-config-diff"); + } + }); + + return () => { + cancelled = true; + }; + }, [selectedHistory?.optimization_job_id, fromVersionKey, toVersionKey]); + + function updateHistoryFilter(key: K, value: AgentHistoryFilters[K]): void { + setHistoryFilters((current) => ({ ...current, [key]: value })); + } + + async function applyHistoryFilters(): Promise { + await refreshHistoryJobs(historyFilters); + } + + function resetHistoryFilters(): void { + const next: AgentHistoryFilters = { status: "all", objective: "all" }; + setHistoryFilters(next); + void refreshHistoryJobs(next).catch((error: unknown) => notifyError(error, "agent-history-filter-reset")); + } + const globalResults = metrics?.global_results || {}; const clientResults = metrics?.client_results || {}; const actualDataLength = globalResults.global_accuracy?.length || 0; @@ -137,7 +258,7 @@ export function AgentResultsCompare(props: AgentPageProps) { const clientTestLossSeries = clientIds.map((cId, idx) => ({ key: `${cId}_test_loss`, label: cId, color: CHART_COLORS[(idx + 1) % CHART_COLORS.length], values: safeSlice(clientResults[cId].test_loss) })); return ( -
+
@@ -147,6 +268,57 @@ export function AgentResultsCompare(props: AgentPageProps) { Past agent optimizations +
+
+ + updateHistoryFilter("q", event.target.value)} + placeholder="Search goal or name" + className="h-8 text-xs" + /> +
+
+ + +
+
+ updateHistoryFilter("dataset", event.target.value)} placeholder="Dataset" className="h-8 text-xs" /> + updateHistoryFilter("config_model", event.target.value)} placeholder="Model" className="h-8 text-xs" /> + updateHistoryFilter("aggregation", event.target.value)} placeholder="Aggregation" className="h-8 text-xs" /> + updateHistoryFilter("model_name", event.target.value)} placeholder="LLM model" className="h-8 text-xs" /> + updateHistoryFilter("num_clients", event.target.value)} placeholder="Clients" className="h-8 text-xs" /> + updateHistoryFilter("num_rounds", event.target.value)} placeholder="Rounds" className="h-8 text-xs" /> + updateHistoryFilter("best_score_min", event.target.value)} placeholder="Min score" className="h-8 text-xs" /> + updateHistoryFilter("best_score_max", event.target.value)} placeholder="Max score" className="h-8 text-xs" /> +
+
+ + +
+
{jobs.length === 0 &&
No history yet.
} {jobs.map((job: any) => { @@ -292,6 +464,91 @@ export function AgentResultsCompare(props: AgentPageProps) {
+ + + + Configuration Versions + + Captured plan and experiment configs + + + {configVersions.length === 0 ? ( +
+ No config versions for this job. +
+ ) : ( + <> +
+ + +
+
+ {configVersions.slice(-3).map((version) => ( +
+
+ {version.label} + {version.source} +
+
+ Iteration {version.iteration || "-"} - {fmt(version.created_at)} +
+
+ ))} +
+
+ {configDiff.length === 0 ? ( +
No config changes.
+ ) : ( +
+ {configDiff.map((change, index) => ( +
+
{change.path}
+ + {change.change_type} + +
+
+ Before + {formatDiffValue(change.old_value)} +
+
+ After + {formatDiffValue(change.new_value)} +
+
+
+ ))} +
+ )} +
+ + )} +
+
+

diff --git a/apps/frontend/src/features/agent/useAgentController.ts b/apps/frontend/src/features/agent/useAgentController.ts index 9ba4108..6234d75 100644 --- a/apps/frontend/src/features/agent/useAgentController.ts +++ b/apps/frontend/src/features/agent/useAgentController.ts @@ -4,6 +4,7 @@ import { toast } from "sonner"; import { agentApi, type AgentExperimentResponse, + type AgentHistoryFilters, type AgentOptimizationObjective, type AgentOptimizationJobSummary, type AgentOptimizeProgressResponse, @@ -86,6 +87,7 @@ export function useAgentController(): AgentPageProps { const [experimentRuns, setExperimentRuns] = useState([]); const [selectedExperimentId, setSelectedExperimentId] = useState(null); const [historyJobs, setHistoryJobs] = useState([]); + const [historyFilters, setHistoryFilters] = useState({ status: "all", objective: "all" }); const [progress, setProgress] = useState(null); const [result, setResult] = useState(null); const [selectedHistory, setSelectedHistory] = useState(null); @@ -108,8 +110,8 @@ export function useAgentController(): AgentPageProps { }); } - async function refreshHistoryJobs(): Promise { - const jobs = await agentApi.listOptimizationJobs(); + async function refreshHistoryJobs(filters: AgentHistoryFilters = historyFilters): Promise { + const jobs = await agentApi.listOptimizationJobs(filters); setHistoryJobs(jobs); if (jobs.length > 0 && selectedHistoryJobId === null && !busy && !progress) { const first = jobs[0]; @@ -411,6 +413,7 @@ export function useAgentController(): AgentPageProps { experimentRuns, goal, handleOptimize, + historyFilters, historyJobs, jobName, lastSubmittedGoal, @@ -427,12 +430,14 @@ export function useAgentController(): AgentPageProps { selectedHistoryJobId, selectExperiment, selectHistoryJob, + refreshHistoryJobs, setConfigConstraint, setGoal, setJobName, setMaxIterations, setModelName, setObjective, + setHistoryFilters, workflowStep, setWorkflowStep, draftPlan, diff --git a/apps/frontend/src/pages/types.ts b/apps/frontend/src/pages/types.ts index dcc0038..ddf4d01 100644 --- a/apps/frontend/src/pages/types.ts +++ b/apps/frontend/src/pages/types.ts @@ -10,6 +10,7 @@ import type { import type { DistributedClient, DistributedJob, DistributedSession, DistributedSessionProgress } from "../api/distributed"; import type { AgentExperimentResponse, + AgentHistoryFilters, AgentOptimizationJobSummary, AgentOptimizeProgressResponse, AgentOptimizeResponse, @@ -241,6 +242,7 @@ export type AgentPageProps = { experimentRuns: AgentRunResponse[]; goal: string; handleOptimize: () => Promise; + historyFilters: AgentHistoryFilters; historyJobs: AgentOptimizationJobSummary[]; jobName: string; lastSubmittedGoal: string | null; @@ -257,6 +259,7 @@ export type AgentPageProps = { selectedHistoryJobId: number | null; selectExperiment: (experimentId: number) => Promise; selectHistoryJob: (optimizationJobId: number) => Promise; + refreshHistoryJobs: (filters?: AgentHistoryFilters) => Promise; setConfigConstraint: (path: string, value: unknown) => void; clearConfigConstraint: (path: string) => void; setGoal: Dispatch>; @@ -264,6 +267,7 @@ export type AgentPageProps = { setMaxIterations: Dispatch>; setModelName: Dispatch>; setObjective: Dispatch>; + setHistoryFilters: Dispatch>; workflowStep: AgentWorkflowStep; setWorkflowStep: Dispatch>; draftPlan: AgentPlanDraft | null; From d9e52d37b64be6e1c1ea4b1305cefefb1196906a Mon Sep 17 00:00:00 2001 From: Fennel1 <627235787@qq.com> Date: Thu, 14 May 2026 22:55:48 +0800 Subject: [PATCH 3/9] update Configuration Versions --- README.md | 12 +- README.zh-CN.md | 12 +- .../agent/optimization_job_repository.py | 4 +- .../app/services/agent/history_service.py | 41 +----- .../services/agent/test_history_service.py | 19 +++ .../agent/components/AgentResultsCompare.tsx | 118 +++++++++++++----- 6 files changed, 120 insertions(+), 86 deletions(-) create mode 100644 apps/backend/tests/services/agent/test_history_service.py diff --git a/README.md b/README.md index d6d94d2..5f6d4f3 100644 --- a/README.md +++ b/README.md @@ -212,12 +212,12 @@ PRs welcome! Figaro is meant to be a readable, research-friendly FL platform. - [x] **Strict Configuration Engine** — Implement strict Pydantic/JSON Schema validation to resolve historical inconsistencies between `config_schema` and underlying algorithms. - [x] **Advanced Experiment Tracking** — Multi-dimensional search filtering (by metrics, hyperparameters, status) and configuration version control (diffing). -**Phase 2: LLM & LoRA Federated Fine-Tuning** -- [ ] **Native LLM Ecosystem Integration** — Seamless Hugging Face model loading (e.g., Llama 3, Qwen) and efficient parsing of JSONL instruction-tuning datasets. -- [ ] **Parameter-Efficient Runtime** — Deep integration with LoRA/PEFT, including support for QLoRA (4-bit/8-bit quantization) to lower client-side memory barriers. -- [ ] **Specialized Adapter Aggregation** — Custom aggregation mechanisms for LoRA adapters, exploring support for heterogeneous LoRA ranks across clients. -- [ ] **LLM Evaluation Metrics** — Built-in evaluation for generative tasks (Rouge, BLEU, Perplexity) and automated LLM-as-a-Judge capabilities. -- [ ] **Hardware Guardrails** — Pre-run dynamic GPU memory estimation (OOM prevention) and automated tuning of gradient accumulation and checkpointing. +**Phase 2: LLM Federated PEFT Fine-Tuning** +- [ ] **LLM/PEFT Configuration Surface** — Extend the existing `config_schema`, Agent planner, and compatibility checks with LLM task type, base model, tokenizer, prompt template, dataset format, and adapter hyperparameters. +- [ ] **Hugging Face Model & Dataset Adapters** — Add adapters behind the current `ModelManager` and data loading pipeline for causal LM / instruction tuning models and JSONL-style supervised fine-tuning datasets. +- [ ] **PEFT Training Runtime** — Add a LoRA-first training path to the existing simulation/distributed runners, with QLoRA-ready quantization options and per-client memory/device controls. +- [ ] **Adapter-Only Federated Aggregation** — Aggregate and persist PEFT adapter weights instead of full model checkpoints, including metadata needed to replay or resume each federated round. +- [ ] **LLM Evaluation & Result Tracking** — Track LLM-specific metrics such as training loss, validation loss, perplexity, token throughput, and adapter checkpoint lineage in the existing Agent results views. **Phase 3: Enterprise & Team Collaboration** - [ ] **Multi-Tenant Workspaces** — Isolated project environments with Role-Based Access Control (RBAC) and comprehensive audit logging. diff --git a/README.zh-CN.md b/README.zh-CN.md index a365f79..864863b 100644 --- a/README.zh-CN.md +++ b/README.zh-CN.md @@ -213,12 +213,12 @@ figaro/ - [x] **配置引擎重构** —— 引入基于 Pydantic/JSON Schema 的严格强校验,彻底修复 `config_schema` 与底层算法实现不一致的问题。 - [x] **高阶实验管理** —— 支持按指标、超参等多维度搜索过滤实验历史,支持配置文件版本控制与 Diff 差异对比。 -**Phase 2:LLM / LoRA 联邦微调支持** -- [ ] **大模型生态原生接入** —— 内置 Hugging Face 适配层,一键加载主流开源模型,支持 JSONL 格式的指令微调数据集高效解析。 -- [ ] **高效分布式微调** —— 深度集成 LoRA/PEFT 训练环境,支持 QLoRA (4-bit/8-bit 量化) 以显著降低边缘节点的显存门槛。 -- [ ] **专属权重聚合策略** —— 针对 LLM 微调定制的 Adapter 权重聚合机制,探索支持客户端异构 LoRA Rank 的聚合方案。 -- [ ] **大模型专项评估体系** —— 集成生成式 NLP 指标,并引入基于大模型的自动化指令跟随能力评估 (LLM-as-a-Judge)。 -- [ ] **资源护栏与预判** —— 训练任务启动前进行动态 GPU 显存预估(防 OOM 机制),并支持梯度累积与 Checkpointing 的自动调优。 +**Phase 2:基于现有框架的 LLM 联邦 PEFT 微调支持** +- [ ] **LLM/PEFT 配置面扩展** —— 在现有 `config_schema`、Agent 规划与兼容性校验中加入任务类型、基础模型、Tokenizer、Prompt 模板、数据格式和 Adapter 超参。 +- [ ] **Hugging Face 模型与数据适配** —— 在当前 `ModelManager` 与数据加载管线后扩展 Causal LM / 指令微调模型适配层,支持 JSONL 等 SFT 数据格式解析。 +- [ ] **PEFT 微调运行时** —— 在现有单机仿真与分布式 Runner 中加入 LoRA 优先的训练路径,并预留 QLoRA 量化选项与客户端显存/设备控制。 +- [ ] **Adapter 权重联邦聚合** —— 聚合和持久化 PEFT Adapter 权重,而不是全量模型 checkpoint,并记录可复现/续跑每轮训练所需的元数据。 +- [ ] **LLM 指标与结果追踪** —— 在现有 Agent 结果视图中追踪 training loss、validation loss、perplexity、token throughput 与 Adapter checkpoint lineage 等指标。 **Phase 3:平台化与企业级协作** - [ ] **多租户与细粒度权限** —— 构建多用户隔离的项目空间 (Workspaces),引入基于角色的访问控制 (RBAC) 和完整的操作审计日志。 diff --git a/apps/backend/app/repositories/agent/optimization_job_repository.py b/apps/backend/app/repositories/agent/optimization_job_repository.py index b014580..7953833 100644 --- a/apps/backend/app/repositories/agent/optimization_job_repository.py +++ b/apps/backend/app/repositories/agent/optimization_job_repository.py @@ -198,11 +198,13 @@ async def find_version( async def get_version(self, version_id: int) -> AgentConfigVersion | None: return await self.session.get(AgentConfigVersion, version_id) - async def list_versions(self, optimization_job_id: int) -> list[AgentConfigVersion]: + async def list_versions(self, optimization_job_id: int, *, source: str | None = None) -> list[AgentConfigVersion]: stmt = ( select(AgentConfigVersion) .where(AgentConfigVersion.optimization_job_id == optimization_job_id) .order_by(AgentConfigVersion.created_at.asc(), AgentConfigVersion.id.asc()) ) + if source is not None: + stmt = stmt.where(AgentConfigVersion.source == source) result = await self.session.execute(stmt) return list(result.scalars().all()) diff --git a/apps/backend/app/services/agent/history_service.py b/apps/backend/app/services/agent/history_service.py index 686a1a2..bbc8132 100644 --- a/apps/backend/app/services/agent/history_service.py +++ b/apps/backend/app/services/agent/history_service.py @@ -173,7 +173,7 @@ async def _ensure_unique_job_name(self, job_name: str, *, exclude_job_id: int | async def list_config_versions(self, optimization_job_id: int) -> list[AgentConfigVersion]: job = await self.get_job_or_raise(optimization_job_id) - versions = await self.config_version_repository.list_versions(optimization_job_id) + versions = await self.config_version_repository.list_versions(optimization_job_id, source="experiment") if versions: return versions snapshot = job.snapshot_json if isinstance(job.snapshot_json, dict) else {} @@ -181,7 +181,7 @@ async def list_config_versions(self, optimization_job_id: int) -> list[AgentConf return [] await self._sync_config_versions(job, snapshot) await self.session.commit() - return await self.config_version_repository.list_versions(optimization_job_id) + return await self.config_version_repository.list_versions(optimization_job_id, source="experiment") async def get_config_version_or_raise(self, optimization_job_id: int, version_id: int) -> AgentConfigVersion: await self.get_job_or_raise(optimization_job_id) @@ -234,30 +234,8 @@ def _build_config_version_candidates(cls, snapshot: dict[str, Any]) -> list[dict candidates: list[dict[str, Any]] = [] previous_config: dict[str, Any] | None = None - draft_experiments = snapshot.get("draft_experiments") - if isinstance(draft_experiments, list): - for index, item in enumerate(draft_experiments, start=1): - if not isinstance(item, dict): - continue - config = item.get("config_patch") or item.get("config") - if not isinstance(config, dict): - continue - candidates.append( - cls._build_config_version_candidate( - config=config, - previous_config=previous_config, - iteration=int(item.get("iteration") or index), - source="draft", - run_id=None, - label=str(item.get("name") or f"Draft {index}"), - diff_json=item.get("config_diff"), - ) - ) - previous_config = config - experiments = snapshot.get("experiments") if isinstance(experiments, list): - previous_config = None for index, item in enumerate(experiments, start=1): if not isinstance(item, dict): continue @@ -277,21 +255,6 @@ def _build_config_version_candidates(cls, snapshot: dict[str, Any]) -> list[dict ) previous_config = config - best_config = snapshot.get("best_config") - if isinstance(best_config, dict): - best_record = cls._best_snapshot_experiment(snapshot) - candidates.append( - cls._build_config_version_candidate( - config=best_config, - previous_config=previous_config, - iteration=int(best_record.get("iteration") or 0) if best_record else 0, - source="best", - run_id=str(best_record.get("run_id")) if best_record and best_record.get("run_id") is not None else None, - label="Best configuration", - diff_json=None, - ) - ) - return candidates @classmethod diff --git a/apps/backend/tests/services/agent/test_history_service.py b/apps/backend/tests/services/agent/test_history_service.py new file mode 100644 index 0000000..5a86db6 --- /dev/null +++ b/apps/backend/tests/services/agent/test_history_service.py @@ -0,0 +1,19 @@ +from app.services.agent.history_service import AgentOptimizationHistoryService + + +def test_config_version_candidates_only_include_experiments(): + snapshot = { + "draft_experiments": [ + {"name": "draft alpha", "config_patch": {"dataset": {"alpha": 0.2}}}, + ], + "experiments": [ + {"run_id": "run-1", "iteration": 1, "name": "alpha_0.1", "config": {"dataset": {"alpha": 0.1}}}, + {"run_id": "run-2", "iteration": 2, "name": "alpha_0.3", "config": {"dataset": {"alpha": 0.3}}}, + ], + "best_config": {"dataset": {"alpha": 0.3}}, + } + + candidates = AgentOptimizationHistoryService._build_config_version_candidates(snapshot) + + assert [candidate["source"] for candidate in candidates] == ["experiment", "experiment"] + assert [candidate["label"] for candidate in candidates] == ["alpha_0.1", "alpha_0.3"] diff --git a/apps/frontend/src/features/agent/components/AgentResultsCompare.tsx b/apps/frontend/src/features/agent/components/AgentResultsCompare.tsx index 6be4ca7..325aaf5 100644 --- a/apps/frontend/src/features/agent/components/AgentResultsCompare.tsx +++ b/apps/frontend/src/features/agent/components/AgentResultsCompare.tsx @@ -3,7 +3,6 @@ import { History, Bot, Activity, Trophy, ArrowRight, Target, Calendar, FileJson, import { Card, CardContent, CardHeader, CardTitle, CardDescription } from "../../../components/ui/card"; import { Badge } from "../../../components/ui/badge"; import { Button } from "../../../components/ui/button"; -import { Input } from "../../../components/ui/input"; import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "../../../components/ui/select"; import { Separator } from "../../../components/ui/separator"; import { MiniLineChart } from "../../simulation/components/MiniLineChart"; @@ -14,7 +13,7 @@ import { baseUrl } from "../../../api/client"; import { agentApi, type AgentConfigChange, type AgentConfigVersion, type AgentHistoryFilters } from "../../../api/agent"; import ReactMarkdown from 'react-markdown'; import remarkGfm from 'remark-gfm'; -import { collectAgentSchemaFields, formatFieldValue, optionLabel, type AgentSchemaField } from "../schema"; +import { collectAgentSchemaFields, formatFieldValue, optionDisabled, optionLabel, type AgentSchemaField } from "../schema"; const CHART_COLORS = [ "hsl(var(--primary))", "#10b981", "#f59e0b", "#ef4444", "#8b5cf6", "#ec4899", "#06b6d4" @@ -53,6 +52,11 @@ const OBJECTIVE_OPTIONS = [ { value: "accuracy", label: "Accuracy" }, ]; +type SelectOption = { + value: string; + label: string; +}; + function asConfigRecord(value: unknown): Record | null { return isRecord(value) ? value : null; } @@ -110,6 +114,40 @@ function diffBadgeClass(changeType: string): string { return "border-amber-500/70 text-amber-600"; } +function schemaSelectOptions(schema: Record | null, path: string): SelectOption[] { + const field = collectAgentSchemaFields(schema, { featuredOnly: false }).find((item) => item.path === path); + if (!field) return []; + return field.options + .filter((option) => !optionDisabled(field.definition, option)) + .map((option) => ({ value: String(option), label: optionLabel(field.definition, option) })); +} + +function optionSelect( + label: string, + value: string | undefined, + options: SelectOption[], + onChange: (value: string) => void, +) { + return ( +
+
{label}
+ +
+ ); +} + export function AgentResultsCompare(props: AgentPageProps) { const { historyFilters, @@ -120,6 +158,7 @@ export function AgentResultsCompare(props: AgentPageProps) { refreshHistoryJobs, setWorkflowStep, configSchema, + modelOptions, notifyError, } = props; @@ -127,7 +166,7 @@ export function AgentResultsCompare(props: AgentPageProps) { const [selectedRunId, setSelectedRunId] = useState(null); const [metrics, setMetrics] = useState(null); const [configVersions, setConfigVersions] = useState([]); - const [fromVersionKey, setFromVersionKey] = useState("baseline"); + const [fromVersionKey, setFromVersionKey] = useState(""); const [toVersionKey, setToVersionKey] = useState(""); const [configDiff, setConfigDiff] = useState([]); @@ -143,6 +182,13 @@ export function AgentResultsCompare(props: AgentPageProps) { () => (bestConfig ? JSON.stringify(bestConfig, null, 2) : ""), [bestConfig], ); + const datasetOptions = useMemo(() => schemaSelectOptions(configSchema, "dataset.name"), [configSchema]); + const configModelOptions = useMemo(() => schemaSelectOptions(configSchema, "model.name"), [configSchema]); + const aggregationOptions = useMemo(() => schemaSelectOptions(configSchema, "federated.aggregation"), [configSchema]); + const llmModelOptions = useMemo( + () => modelOptions.map((item) => ({ value: item, label: item })), + [modelOptions], + ); useEffect(() => { if (bestExp) { @@ -170,7 +216,7 @@ export function AgentResultsCompare(props: AgentPageProps) { const optimizationJobId = selectedHistory?.optimization_job_id; if (!optimizationJobId) { setConfigVersions([]); - setFromVersionKey("baseline"); + setFromVersionKey(""); setToVersionKey(""); setConfigDiff([]); return; @@ -180,13 +226,19 @@ export function AgentResultsCompare(props: AgentPageProps) { agentApi.listConfigVersions(optimizationJobId) .then((versions) => { if (cancelled) return; - setConfigVersions(versions); - setFromVersionKey(versions.length > 1 ? String(versions[0].id) : "baseline"); - setToVersionKey(versions.length > 0 ? String(versions[versions.length - 1].id) : ""); + const experimentVersions = versions.filter((version) => version.source === "experiment"); + const firstVersion = experimentVersions[0]; + const lastVersion = experimentVersions[experimentVersions.length - 1]; + setConfigVersions(experimentVersions); + setFromVersionKey(firstVersion ? String(firstVersion.id) : ""); + setToVersionKey(lastVersion ? String(lastVersion.id) : ""); + setConfigDiff([]); }) .catch((error) => { if (!cancelled) { setConfigVersions([]); + setFromVersionKey(""); + setToVersionKey(""); setConfigDiff([]); notifyError(error, "agent-config-versions"); } @@ -199,18 +251,25 @@ export function AgentResultsCompare(props: AgentPageProps) { useEffect(() => { const optimizationJobId = selectedHistory?.optimization_job_id; + const fromVersionId = Number(fromVersionKey); const toVersionId = Number(toVersionKey); - if (!optimizationJobId || !toVersionKey || !Number.isFinite(toVersionId)) { + if ( + !optimizationJobId || + !fromVersionKey || + !toVersionKey || + !Number.isFinite(fromVersionId) || + !Number.isFinite(toVersionId) || + fromVersionId === toVersionId + ) { setConfigDiff([]); return; } - const fromVersionId = fromVersionKey === "baseline" ? null : Number(fromVersionKey); let cancelled = false; agentApi.getConfigDiff( optimizationJobId, toVersionId, - fromVersionId !== null && Number.isFinite(fromVersionId) && fromVersionId !== toVersionId ? fromVersionId : null, + fromVersionId, ) .then((payload) => { if (!cancelled) { @@ -234,7 +293,8 @@ export function AgentResultsCompare(props: AgentPageProps) { } async function applyHistoryFilters(): Promise { - await refreshHistoryJobs(historyFilters); + const { q: _unusedTextFilter, ...selectableFilters } = historyFilters; + await refreshHistoryJobs(selectableFilters); } function resetHistoryFilters(): void { @@ -268,15 +328,10 @@ export function AgentResultsCompare(props: AgentPageProps) { Past agent optimizations -
-
+
+
- updateHistoryFilter("q", event.target.value)} - placeholder="Search goal or name" - className="h-8 text-xs" - /> + Filter history
- updateHistoryFilter("dataset", event.target.value)} placeholder="Dataset" className="h-8 text-xs" /> - updateHistoryFilter("config_model", event.target.value)} placeholder="Model" className="h-8 text-xs" /> - updateHistoryFilter("aggregation", event.target.value)} placeholder="Aggregation" className="h-8 text-xs" /> - updateHistoryFilter("model_name", event.target.value)} placeholder="LLM model" className="h-8 text-xs" /> - updateHistoryFilter("num_clients", event.target.value)} placeholder="Clients" className="h-8 text-xs" /> - updateHistoryFilter("num_rounds", event.target.value)} placeholder="Rounds" className="h-8 text-xs" /> - updateHistoryFilter("best_score_min", event.target.value)} placeholder="Min score" className="h-8 text-xs" /> - updateHistoryFilter("best_score_max", event.target.value)} placeholder="Max score" className="h-8 text-xs" /> + {optionSelect("Dataset", historyFilters.dataset, datasetOptions, (value) => updateHistoryFilter("dataset", value))} + {optionSelect("Model", historyFilters.config_model, configModelOptions, (value) => updateHistoryFilter("config_model", value))} + {optionSelect("Aggregation", historyFilters.aggregation, aggregationOptions, (value) => updateHistoryFilter("aggregation", value))} + {optionSelect("LLM", historyFilters.model_name, llmModelOptions, (value) => updateHistoryFilter("model_name", value))}