diff --git a/src/gausse/optim/sweep.py b/src/gausse/optim/sweep.py new file mode 100644 index 0000000..60f60d3 --- /dev/null +++ b/src/gausse/optim/sweep.py @@ -0,0 +1,105 @@ +"""Массовый Monte Carlo sweep: честный широкий перебор, все прогоны в SQLite. + +Воркеры (пул процессов) только сэмплируют и оценивают геномы; запись в +SQLite идёт через единственный writer-процесс поверх очереди +(`storage.database.run_writer_process`), чтобы не было конкурентной записи +в один файл базы. +""" + +import json +import multiprocessing +import random +import uuid +from concurrent.futures import ProcessPoolExecutor +from datetime import datetime, timezone +from pathlib import Path +from typing import Callable + +from gausse.components.database import ComponentDatabase +from gausse.optim.objective import decoded_summary, evaluate +from gausse.optim.search_space import SearchBounds, genome_to_dict, sample_genome +from gausse.storage.database import run_writer_process +from gausse.storage.schema import RunRecord + +MODEL_VERSION = "gausse-physics-v1" + +_worker_db: ComponentDatabase | None = None +_worker_bounds: SearchBounds | None = None + + +def _init_worker(data_dir: Path | None, bounds: SearchBounds) -> None: + global _worker_db, _worker_bounds + _worker_db = ComponentDatabase.load(data_dir) if data_dir else ComponentDatabase.load() + _worker_bounds = bounds + + +def _evaluate_one(seed: int, search_mode: str = "sweep") -> RunRecord: + rng = random.Random(seed) + genome = sample_genome(_worker_db, _worker_bounds, rng) + return _evaluate_genome(genome, search_mode) + + +def _evaluate_genome(genome, search_mode: str) -> RunRecord: + result = evaluate(genome, _worker_db, _worker_bounds) + summary = decoded_summary(genome, _worker_db) + return RunRecord( + run_id=str(uuid.uuid4()), + timestamp=datetime.now(timezone.utc).isoformat(), + search_mode=search_mode, + genome_json=json.dumps(genome_to_dict(genome)), + decoded_summary_json=json.dumps(summary, ensure_ascii=False), + feasible=result.feasible, + model_version=MODEL_VERSION, + infeasible_reason=result.reason, + failed_stage_index=result.failed_stage_index, + efficiency=result.efficiency, + exit_velocity_mps=result.exit_velocity_mps, + cost_rub=result.cost_rub, + energy_breakdown_json=json.dumps(result.energy_breakdown) if result.energy_breakdown else None, + ) + + +def run_sweep( + db_path: Path, + n_runs: int, + bounds: SearchBounds = SearchBounds(), + data_dir: Path | None = None, + n_workers: int | None = None, + seed: int | None = None, + progress_callback: Callable[[int, int], None] | None = None, +) -> dict: + """Прогоняет n_runs случайных конфигураций параллельно, пишет каждую в SQLite. + + Возвращает сводку {n_runs, n_feasible} — не для принятия решений (для + этого нужно читать саму базу), а как быстрая сверка "ничего не потерялось". + """ + n_workers = n_workers or multiprocessing.cpu_count() + base_seed = seed if seed is not None else random.randrange(2**31) + seeds = [base_seed + i for i in range(n_runs)] + + ctx = multiprocessing.get_context("spawn") + queue = ctx.Queue() + writer = ctx.Process(target=run_writer_process, args=(queue, db_path)) + writer.start() + + completed = 0 + n_feasible = 0 + try: + with ProcessPoolExecutor( + max_workers=n_workers, + mp_context=ctx, + initializer=_init_worker, + initargs=(data_dir, bounds), + ) as executor: + for record in executor.map(_evaluate_one, seeds): + queue.put(record) + completed += 1 + if record.feasible: + n_feasible += 1 + if progress_callback: + progress_callback(completed, n_runs) + finally: + queue.put(None) + writer.join() + + return {"n_runs": completed, "n_feasible": n_feasible} diff --git a/tests/test_sweep.py b/tests/test_sweep.py new file mode 100644 index 0000000..d002bbb --- /dev/null +++ b/tests/test_sweep.py @@ -0,0 +1,53 @@ +import json + +from gausse.optim.search_space import SearchBounds, genome_from_dict +from gausse.optim.sweep import MODEL_VERSION, run_sweep +from gausse.storage.database import count_runs, fetch_runs, open_connection + + +def test_sweep_runs_all_seeds_and_records_everything(tmp_path): + db_path = tmp_path / "runs.sqlite3" + bounds = SearchBounds(max_stages=2) # маленькое пространство -> быстрый тест + + summary = run_sweep(db_path, n_runs=12, bounds=bounds, n_workers=2, seed=123) + + assert summary["n_runs"] == 12 + conn = open_connection(db_path) + assert count_runs(conn) == 12 + # хотя бы какие-то честно попали и в успех, и в провал -- иначе тест + # пространства слишком узкий/широкий, чтобы быть показательным + rows = fetch_runs(conn) + assert len(rows) == 12 + for row in rows: + assert row.search_mode == "sweep" + assert row.model_version == MODEL_VERSION + # genome_json должен быть валидным и декодируемым геномом + genome = genome_from_dict(json.loads(row.genome_json)) + assert len(genome.stages) >= 1 + summary_dict = json.loads(row.decoded_summary_json) + assert "stages" in summary_dict + if row.feasible: + assert row.efficiency is not None + assert row.infeasible_reason is None + else: + assert row.infeasible_reason is not None + assert row.efficiency is None + conn.close() + + +def test_sweep_is_deterministic_given_same_seed(tmp_path): + bounds = SearchBounds(max_stages=2) + db_path_a = tmp_path / "a.sqlite3" + db_path_b = tmp_path / "b.sqlite3" + + run_sweep(db_path_a, n_runs=8, bounds=bounds, n_workers=1, seed=42) + run_sweep(db_path_b, n_runs=8, bounds=bounds, n_workers=1, seed=42) + + conn_a = open_connection(db_path_a) + conn_b = open_connection(db_path_b) + rows_a = sorted(fetch_runs(conn_a), key=lambda r: r.genome_json) + rows_b = sorted(fetch_runs(conn_b), key=lambda r: r.genome_json) + assert [r.genome_json for r in rows_a] == [r.genome_json for r in rows_b] + assert [r.efficiency for r in rows_a] == [r.efficiency for r in rows_b] + conn_a.close() + conn_b.close()