asi-evolve 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- asi_evolve/__init__.py +22 -0
- asi_evolve/_vendor/ASI_Evolve/__init__.py +3 -0
- asi_evolve/_vendor/ASI_Evolve/assets/Overview.png +0 -0
- asi_evolve/_vendor/ASI_Evolve/cognition/__init__.py +7 -0
- asi_evolve/_vendor/ASI_Evolve/cognition/cognition.py +180 -0
- asi_evolve/_vendor/ASI_Evolve/config.yaml +119 -0
- asi_evolve/_vendor/ASI_Evolve/database/__init__.py +21 -0
- asi_evolve/_vendor/ASI_Evolve/database/algorithms/__init__.py +17 -0
- asi_evolve/_vendor/ASI_Evolve/database/algorithms/base.py +43 -0
- asi_evolve/_vendor/ASI_Evolve/database/algorithms/factory.py +41 -0
- asi_evolve/_vendor/ASI_Evolve/database/algorithms/greedy.py +18 -0
- asi_evolve/_vendor/ASI_Evolve/database/algorithms/island.py +607 -0
- asi_evolve/_vendor/ASI_Evolve/database/algorithms/random.py +19 -0
- asi_evolve/_vendor/ASI_Evolve/database/algorithms/ucb1.py +62 -0
- asi_evolve/_vendor/ASI_Evolve/database/database.py +276 -0
- asi_evolve/_vendor/ASI_Evolve/database/embedding.py +64 -0
- asi_evolve/_vendor/ASI_Evolve/database/faiss_index.py +193 -0
- asi_evolve/_vendor/ASI_Evolve/main.py +110 -0
- asi_evolve/_vendor/ASI_Evolve/pipeline/__init__.py +7 -0
- asi_evolve/_vendor/ASI_Evolve/pipeline/analyzer/__init__.py +6 -0
- asi_evolve/_vendor/ASI_Evolve/pipeline/analyzer/analyzer.py +60 -0
- asi_evolve/_vendor/ASI_Evolve/pipeline/base.py +53 -0
- asi_evolve/_vendor/ASI_Evolve/pipeline/engineer/__init__.py +6 -0
- asi_evolve/_vendor/ASI_Evolve/pipeline/engineer/engineer.py +264 -0
- asi_evolve/_vendor/ASI_Evolve/pipeline/main.py +572 -0
- asi_evolve/_vendor/ASI_Evolve/pipeline/manager/__init__.py +6 -0
- asi_evolve/_vendor/ASI_Evolve/pipeline/manager/manager.py +47 -0
- asi_evolve/_vendor/ASI_Evolve/pipeline/researcher/__init__.py +6 -0
- asi_evolve/_vendor/ASI_Evolve/pipeline/researcher/researcher.py +165 -0
- asi_evolve/_vendor/ASI_Evolve/skills/evolve/agents/openai.yaml +7 -0
- asi_evolve/_vendor/ASI_Evolve/utils/__init__.py +23 -0
- asi_evolve/_vendor/ASI_Evolve/utils/best_snapshot.py +66 -0
- asi_evolve/_vendor/ASI_Evolve/utils/config.py +103 -0
- asi_evolve/_vendor/ASI_Evolve/utils/diff.py +140 -0
- asi_evolve/_vendor/ASI_Evolve/utils/llm.py +306 -0
- asi_evolve/_vendor/ASI_Evolve/utils/logger.py +228 -0
- asi_evolve/_vendor/ASI_Evolve/utils/prompt.py +155 -0
- asi_evolve/_vendor/ASI_Evolve/utils/prompts/analyzer.jinja2 +54 -0
- asi_evolve/_vendor/ASI_Evolve/utils/prompts/manager.jinja2 +15 -0
- asi_evolve/_vendor/ASI_Evolve/utils/prompts/researcher.jinja2 +33 -0
- asi_evolve/_vendor/ASI_Evolve/utils/prompts/researcher_diff.jinja2 +89 -0
- asi_evolve/_vendor/ASI_Evolve/utils/structures.py +145 -0
- asi_evolve/cli.py +311 -0
- asi_evolve/providers/__init__.py +202 -0
- asi_evolve/runner.py +642 -0
- asi_evolve/serve.py +180 -0
- asi_evolve/static/css/dashboard.css +147 -0
- asi_evolve/static/js/dashboard.js +104 -0
- asi_evolve/templates/index.html +58 -0
- asi_evolve-0.1.0.dist-info/METADATA +1653 -0
- asi_evolve-0.1.0.dist-info/RECORD +57 -0
- asi_evolve-0.1.0.dist-info/WHEEL +5 -0
- asi_evolve-0.1.0.dist-info/entry_points.txt +2 -0
- asi_evolve-0.1.0.dist-info/licenses/LICENSE +35 -0
- asi_evolve-0.1.0.dist-info/licenses/LICENSE-VENDORED +201 -0
- asi_evolve-0.1.0.dist-info/licenses/NOTICE +54 -0
- asi_evolve-0.1.0.dist-info/top_level.txt +1 -0
asi_evolve/__init__.py
ADDED
|
@@ -0,0 +1,22 @@
|
|
|
1
|
+
"""asi-evolve — Pythonic wrapper around the ASI-Evolve evolutionary-search framework.
|
|
2
|
+
|
|
3
|
+
Drop a problem, an initial candidate, and an evaluator script. asi-evolve runs the
|
|
4
|
+
autonomous loop until it finds a better solution. CLI or `--serve` web dashboard.
|
|
5
|
+
|
|
6
|
+
Vendored upstream: GAIR-NLP/ASI-Evolve @ fb8a67e (Apache 2.0).
|
|
7
|
+
This wrapper layer: MIT (see LICENSE).
|
|
8
|
+
"""
|
|
9
|
+
|
|
10
|
+
__version__ = "0.1.0"
|
|
11
|
+
|
|
12
|
+
# IMPORTANT: do not eagerly import the vendored package here — it pulls in
|
|
13
|
+
# faiss-cpu, sentence-transformers, torch, numpy and takes ~20s to load.
|
|
14
|
+
# The `Pipeline` re-export is lazy so `asi-run --help` and pure-Python
|
|
15
|
+
# tooling stay fast.
|
|
16
|
+
def __getattr__(name):
|
|
17
|
+
if name == "Pipeline":
|
|
18
|
+
from asi_evolve._vendor.ASI_Evolve.pipeline import Pipeline as _Pipeline
|
|
19
|
+
return _Pipeline
|
|
20
|
+
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
|
|
21
|
+
|
|
22
|
+
__all__ = ["Pipeline", "__version__"]
|
|
Binary file
|
|
@@ -0,0 +1,180 @@
|
|
|
1
|
+
"""Cognition-store implementation."""
|
|
2
|
+
|
|
3
|
+
import json
|
|
4
|
+
import uuid
|
|
5
|
+
from pathlib import Path
|
|
6
|
+
from threading import RLock
|
|
7
|
+
from typing import Dict, List, Optional, Tuple
|
|
8
|
+
|
|
9
|
+
from ..utils.structures import CognitionItem
|
|
10
|
+
from ..database.faiss_index import FAISSIndex
|
|
11
|
+
from ..database.embedding import EmbeddingService
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
class Cognition:
|
|
15
|
+
"""
|
|
16
|
+
Persistent cognition store with embedding-based retrieval.
|
|
17
|
+
|
|
18
|
+
Features:
|
|
19
|
+
- CRUD operations for cognition items
|
|
20
|
+
- Semantic retrieval through FAISS
|
|
21
|
+
- Local persistence to disk
|
|
22
|
+
"""
|
|
23
|
+
|
|
24
|
+
def __init__(
|
|
25
|
+
self,
|
|
26
|
+
storage_dir: Path,
|
|
27
|
+
embedding_model: str = "sentence-transformers/all-MiniLM-L6-v2",
|
|
28
|
+
embedding_dim: int = 384,
|
|
29
|
+
retrieval_top_k: int = 5,
|
|
30
|
+
score_threshold: float = 0.5,
|
|
31
|
+
faiss_index_type: str = "IP",
|
|
32
|
+
):
|
|
33
|
+
self.storage_dir = Path(storage_dir)
|
|
34
|
+
self.storage_dir.mkdir(parents=True, exist_ok=True)
|
|
35
|
+
|
|
36
|
+
self.lock = RLock()
|
|
37
|
+
self.retrieval_top_k = retrieval_top_k
|
|
38
|
+
self.score_threshold = score_threshold
|
|
39
|
+
|
|
40
|
+
self.items: Dict[str, CognitionItem] = {}
|
|
41
|
+
|
|
42
|
+
self.embedding = EmbeddingService(model_name=embedding_model)
|
|
43
|
+
|
|
44
|
+
self.faiss = FAISSIndex(
|
|
45
|
+
dimension=embedding_dim,
|
|
46
|
+
index_type=faiss_index_type,
|
|
47
|
+
storage_path=self.storage_dir / "faiss",
|
|
48
|
+
)
|
|
49
|
+
|
|
50
|
+
self.str_to_int: Dict[str, int] = {}
|
|
51
|
+
self.int_to_str: Dict[int, str] = {}
|
|
52
|
+
self.next_int_id = 0
|
|
53
|
+
|
|
54
|
+
self._load()
|
|
55
|
+
|
|
56
|
+
def _get_int_id(self, str_id: str) -> int:
|
|
57
|
+
if str_id not in self.str_to_int:
|
|
58
|
+
self.str_to_int[str_id] = self.next_int_id
|
|
59
|
+
self.int_to_str[self.next_int_id] = str_id
|
|
60
|
+
self.next_int_id += 1
|
|
61
|
+
return self.str_to_int[str_id]
|
|
62
|
+
|
|
63
|
+
def add(self, item: CognitionItem) -> str:
|
|
64
|
+
with self.lock:
|
|
65
|
+
if not item.id:
|
|
66
|
+
item.id = str(uuid.uuid4())
|
|
67
|
+
|
|
68
|
+
self.items[item.id] = item
|
|
69
|
+
|
|
70
|
+
if item.content:
|
|
71
|
+
int_id = self._get_int_id(item.id)
|
|
72
|
+
vector = self.embedding.encode(item.content)
|
|
73
|
+
self.faiss.add(int_id, vector)
|
|
74
|
+
|
|
75
|
+
self._save()
|
|
76
|
+
return item.id
|
|
77
|
+
|
|
78
|
+
def add_batch(self, items: List[CognitionItem]) -> List[str]:
|
|
79
|
+
return [self.add(item) for item in items]
|
|
80
|
+
|
|
81
|
+
def remove(self, item_id: str) -> bool:
|
|
82
|
+
with self.lock:
|
|
83
|
+
if item_id not in self.items:
|
|
84
|
+
return False
|
|
85
|
+
|
|
86
|
+
del self.items[item_id]
|
|
87
|
+
|
|
88
|
+
if item_id in self.str_to_int:
|
|
89
|
+
int_id = self.str_to_int.pop(item_id)
|
|
90
|
+
self.int_to_str.pop(int_id, None)
|
|
91
|
+
self.faiss.remove(int_id)
|
|
92
|
+
|
|
93
|
+
self._save()
|
|
94
|
+
return True
|
|
95
|
+
|
|
96
|
+
def remove_batch(self, item_ids: List[str]) -> int:
|
|
97
|
+
return sum(1 for iid in item_ids if self.remove(iid))
|
|
98
|
+
|
|
99
|
+
def get(self, item_id: str) -> Optional[CognitionItem]:
|
|
100
|
+
return self.items.get(item_id)
|
|
101
|
+
|
|
102
|
+
def get_all(self) -> List[CognitionItem]:
|
|
103
|
+
return list(self.items.values())
|
|
104
|
+
|
|
105
|
+
def retrieve(
|
|
106
|
+
self,
|
|
107
|
+
query: str,
|
|
108
|
+
top_k: Optional[int] = None,
|
|
109
|
+
score_threshold: Optional[float] = None,
|
|
110
|
+
) -> List[Tuple[CognitionItem, float]]:
|
|
111
|
+
top_k = top_k or self.retrieval_top_k
|
|
112
|
+
score_threshold = score_threshold if score_threshold is not None else self.score_threshold
|
|
113
|
+
|
|
114
|
+
query_vector = self.embedding.encode(query)
|
|
115
|
+
results = self.faiss.search(query_vector, top_k, score_threshold)
|
|
116
|
+
|
|
117
|
+
items_with_scores = []
|
|
118
|
+
for int_id, score in results:
|
|
119
|
+
str_id = self.int_to_str.get(int_id)
|
|
120
|
+
if str_id:
|
|
121
|
+
item = self.items.get(str_id)
|
|
122
|
+
if item:
|
|
123
|
+
items_with_scores.append((item, score))
|
|
124
|
+
|
|
125
|
+
return items_with_scores
|
|
126
|
+
|
|
127
|
+
def search(self, query: str, top_k: Optional[int] = None) -> List[CognitionItem]:
|
|
128
|
+
results = self.retrieve(query, top_k)
|
|
129
|
+
return [item for item, _ in results]
|
|
130
|
+
|
|
131
|
+
def reset(self):
|
|
132
|
+
with self.lock:
|
|
133
|
+
self.items.clear()
|
|
134
|
+
self.str_to_int.clear()
|
|
135
|
+
self.int_to_str.clear()
|
|
136
|
+
self.next_int_id = 0
|
|
137
|
+
self.faiss.reset()
|
|
138
|
+
|
|
139
|
+
data_file = self.storage_dir / "cognition.json"
|
|
140
|
+
if data_file.exists():
|
|
141
|
+
data_file.unlink()
|
|
142
|
+
|
|
143
|
+
def _save(self):
|
|
144
|
+
data_file = self.storage_dir / "cognition.json"
|
|
145
|
+
|
|
146
|
+
data = {
|
|
147
|
+
"items": {k: v.to_dict() for k, v in self.items.items()},
|
|
148
|
+
"str_to_int": self.str_to_int,
|
|
149
|
+
"int_to_str": {str(k): v for k, v in self.int_to_str.items()},
|
|
150
|
+
"next_int_id": self.next_int_id,
|
|
151
|
+
}
|
|
152
|
+
|
|
153
|
+
with open(data_file, "w", encoding="utf-8") as f:
|
|
154
|
+
json.dump(data, f, ensure_ascii=False, indent=2)
|
|
155
|
+
|
|
156
|
+
self.faiss.save()
|
|
157
|
+
|
|
158
|
+
def _load(self):
|
|
159
|
+
data_file = self.storage_dir / "cognition.json"
|
|
160
|
+
|
|
161
|
+
if not data_file.exists():
|
|
162
|
+
return
|
|
163
|
+
|
|
164
|
+
with open(data_file, "r", encoding="utf-8") as f:
|
|
165
|
+
data = json.load(f)
|
|
166
|
+
|
|
167
|
+
for item_id, item_data in data.get("items", {}).items():
|
|
168
|
+
item = CognitionItem.from_dict(item_data)
|
|
169
|
+
self.items[item_id] = item
|
|
170
|
+
|
|
171
|
+
self.str_to_int = data.get("str_to_int", {})
|
|
172
|
+
self.int_to_str = {int(k): v for k, v in data.get("int_to_str", {}).items()}
|
|
173
|
+
self.next_int_id = data.get("next_int_id", 0)
|
|
174
|
+
|
|
175
|
+
@property
|
|
176
|
+
def size(self) -> int:
|
|
177
|
+
return len(self.items)
|
|
178
|
+
|
|
179
|
+
def __len__(self) -> int:
|
|
180
|
+
return self.size
|
|
@@ -0,0 +1,119 @@
|
|
|
1
|
+
# Evolve Framework Configuration (Default Template)
|
|
2
|
+
# ================================================
|
|
3
|
+
# This file defines repository-wide defaults.
|
|
4
|
+
# Priority: explicit config file > experiment config > repository config.
|
|
5
|
+
|
|
6
|
+
# Experiment name
|
|
7
|
+
experiment_name: "default"
|
|
8
|
+
|
|
9
|
+
# API configuration
|
|
10
|
+
# In practice, experiment-level config files usually override these values.
|
|
11
|
+
api:
|
|
12
|
+
provider: "openai" # Any OpenAI-compatible endpoint
|
|
13
|
+
base_url: "your_base_url"
|
|
14
|
+
api_key: "your_api_key"
|
|
15
|
+
model: "your_model"
|
|
16
|
+
temperature: 0
|
|
17
|
+
top_p: 0.8
|
|
18
|
+
max_tokens: 16384
|
|
19
|
+
seed: 42
|
|
20
|
+
timeout: 300
|
|
21
|
+
retry_times: 3
|
|
22
|
+
retry_delay: 5
|
|
23
|
+
|
|
24
|
+
# Logging configuration
|
|
25
|
+
logging:
|
|
26
|
+
level: "INFO"
|
|
27
|
+
console: true
|
|
28
|
+
wandb:
|
|
29
|
+
enabled: true
|
|
30
|
+
offline: false # Use `wandb sync` later when running offline
|
|
31
|
+
project: "evolve"
|
|
32
|
+
entity: ""
|
|
33
|
+
|
|
34
|
+
# Pipeline configuration
|
|
35
|
+
pipeline:
|
|
36
|
+
# Agent toggles
|
|
37
|
+
agents:
|
|
38
|
+
manager: false
|
|
39
|
+
researcher: true
|
|
40
|
+
engineer: true
|
|
41
|
+
analyzer: true
|
|
42
|
+
|
|
43
|
+
# Researcher configuration
|
|
44
|
+
researcher:
|
|
45
|
+
# When enabled, the researcher edits a sampled parent via SEARCH/REPLACE diffs.
|
|
46
|
+
diff_based_evolution: true
|
|
47
|
+
|
|
48
|
+
# Regex used to parse diff-based edits.
|
|
49
|
+
diff_pattern: "<<<<<<< SEARCH\\n(.*?)=======\\n(.*?)>>>>>>> REPLACE"
|
|
50
|
+
|
|
51
|
+
# Maximum generated code length.
|
|
52
|
+
max_code_length: 10000
|
|
53
|
+
|
|
54
|
+
# Retry counts per agent
|
|
55
|
+
max_retries:
|
|
56
|
+
researcher: 3
|
|
57
|
+
engineer: 2
|
|
58
|
+
analyzer: 2
|
|
59
|
+
|
|
60
|
+
# Engineer timeout in seconds
|
|
61
|
+
engineer_timeout: 1800
|
|
62
|
+
|
|
63
|
+
# Parallel execution
|
|
64
|
+
parallel:
|
|
65
|
+
num_workers: 1
|
|
66
|
+
# Suggested settings:
|
|
67
|
+
# - 1: sequential execution, best for debugging
|
|
68
|
+
# - 2-4: practical parallelism for production runs
|
|
69
|
+
# - >4: requires enough API quota and compute resources
|
|
70
|
+
|
|
71
|
+
# Number of historical nodes sampled per step
|
|
72
|
+
sample_n: 3
|
|
73
|
+
|
|
74
|
+
# Optional LLM-based judge
|
|
75
|
+
# final_score = (1 - judge_ratio) * eval_score + judge_ratio * judge_score
|
|
76
|
+
judge:
|
|
77
|
+
enabled: false
|
|
78
|
+
ratio: 0.2
|
|
79
|
+
# Provide `judge.jinja2` in the experiment prompts directory when enabled.
|
|
80
|
+
|
|
81
|
+
# Cognition configuration
|
|
82
|
+
cognition:
|
|
83
|
+
# Path relative to the experiment directory
|
|
84
|
+
storage_dir: "cognition_data"
|
|
85
|
+
embedding:
|
|
86
|
+
model: "sentence-transformers/all-MiniLM-L6-v2"
|
|
87
|
+
dimension: 384
|
|
88
|
+
faiss:
|
|
89
|
+
index_type: "IP"
|
|
90
|
+
retrieval:
|
|
91
|
+
top_k: 3
|
|
92
|
+
score_threshold: 0.3
|
|
93
|
+
# Reserved for future web-search support
|
|
94
|
+
web_search:
|
|
95
|
+
enabled: false
|
|
96
|
+
|
|
97
|
+
# Database configuration
|
|
98
|
+
database:
|
|
99
|
+
# Path relative to the experiment directory
|
|
100
|
+
storage_dir: "database_data"
|
|
101
|
+
# Limit database size. When full, the lowest-scoring node is removed.
|
|
102
|
+
max_size: null
|
|
103
|
+
embedding:
|
|
104
|
+
model: "sentence-transformers/all-MiniLM-L6-v2"
|
|
105
|
+
dimension: 384
|
|
106
|
+
# Sampling algorithm
|
|
107
|
+
sampling:
|
|
108
|
+
algorithm: "ucb1"
|
|
109
|
+
ucb1_c: 1.414
|
|
110
|
+
# Island-sampler parameters
|
|
111
|
+
island:
|
|
112
|
+
num_islands: 5
|
|
113
|
+
migration_interval: 10
|
|
114
|
+
migration_rate: 0.1
|
|
115
|
+
exploration_ratio: 0.2
|
|
116
|
+
exploitation_ratio: 0.3
|
|
117
|
+
# The remaining probability mass is used for weighted sampling.
|
|
118
|
+
faiss:
|
|
119
|
+
index_type: "IP"
|
|
@@ -0,0 +1,21 @@
|
|
|
1
|
+
"""Experiment database and sampling utilities."""
|
|
2
|
+
|
|
3
|
+
from .database import Database
|
|
4
|
+
from .algorithms import (
|
|
5
|
+
get_sampler,
|
|
6
|
+
BaseSampler,
|
|
7
|
+
UCB1Sampler,
|
|
8
|
+
RandomSampler,
|
|
9
|
+
GreedySampler,
|
|
10
|
+
IslandSampler,
|
|
11
|
+
)
|
|
12
|
+
|
|
13
|
+
__all__ = [
|
|
14
|
+
"Database",
|
|
15
|
+
"get_sampler",
|
|
16
|
+
"BaseSampler",
|
|
17
|
+
"UCB1Sampler",
|
|
18
|
+
"RandomSampler",
|
|
19
|
+
"GreedySampler",
|
|
20
|
+
"IslandSampler",
|
|
21
|
+
]
|
|
@@ -0,0 +1,17 @@
|
|
|
1
|
+
"""Sampling algorithms used by the experiment database."""
|
|
2
|
+
|
|
3
|
+
from .base import BaseSampler
|
|
4
|
+
from .random import RandomSampler
|
|
5
|
+
from .greedy import GreedySampler
|
|
6
|
+
from .ucb1 import UCB1Sampler
|
|
7
|
+
from .island import IslandSampler
|
|
8
|
+
from .factory import get_sampler
|
|
9
|
+
|
|
10
|
+
__all__ = [
|
|
11
|
+
"BaseSampler",
|
|
12
|
+
"RandomSampler",
|
|
13
|
+
"GreedySampler",
|
|
14
|
+
"UCB1Sampler",
|
|
15
|
+
"IslandSampler",
|
|
16
|
+
"get_sampler",
|
|
17
|
+
]
|
|
@@ -0,0 +1,43 @@
|
|
|
1
|
+
"""Base interface for node samplers."""
|
|
2
|
+
|
|
3
|
+
from abc import ABC, abstractmethod
|
|
4
|
+
from typing import List, TYPE_CHECKING
|
|
5
|
+
|
|
6
|
+
if TYPE_CHECKING:
|
|
7
|
+
from ..utils.structures import Node
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
class BaseSampler(ABC):
|
|
11
|
+
"""Abstract sampler used by the experiment database."""
|
|
12
|
+
|
|
13
|
+
@abstractmethod
|
|
14
|
+
def sample(self, nodes: List["Node"], n: int) -> List["Node"]:
|
|
15
|
+
"""
|
|
16
|
+
Sample a subset of nodes.
|
|
17
|
+
|
|
18
|
+
Args:
|
|
19
|
+
nodes: All candidate nodes.
|
|
20
|
+
n: Number of nodes to return.
|
|
21
|
+
|
|
22
|
+
Returns:
|
|
23
|
+
A list of sampled nodes.
|
|
24
|
+
"""
|
|
25
|
+
pass
|
|
26
|
+
|
|
27
|
+
def on_node_added(self, node: "Node") -> None:
|
|
28
|
+
"""
|
|
29
|
+
Hook called when a node is added.
|
|
30
|
+
|
|
31
|
+
Args:
|
|
32
|
+
node: The newly added node.
|
|
33
|
+
"""
|
|
34
|
+
pass
|
|
35
|
+
|
|
36
|
+
def on_node_removed(self, node: "Node") -> None:
|
|
37
|
+
"""
|
|
38
|
+
Hook called when a node is removed.
|
|
39
|
+
|
|
40
|
+
Args:
|
|
41
|
+
node: The removed node.
|
|
42
|
+
"""
|
|
43
|
+
pass
|
|
@@ -0,0 +1,41 @@
|
|
|
1
|
+
"""Factory for sampler implementations."""
|
|
2
|
+
|
|
3
|
+
from typing import Any
|
|
4
|
+
|
|
5
|
+
from .base import BaseSampler
|
|
6
|
+
from .random import RandomSampler
|
|
7
|
+
from .greedy import GreedySampler
|
|
8
|
+
from .ucb1 import UCB1Sampler
|
|
9
|
+
from .island import IslandSampler
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
def get_sampler(algorithm: str, **kwargs) -> BaseSampler:
|
|
13
|
+
"""
|
|
14
|
+
Construct a sampler by name.
|
|
15
|
+
|
|
16
|
+
Args:
|
|
17
|
+
algorithm: Algorithm name (`ucb1`, `random`, `greedy`, or `island`).
|
|
18
|
+
**kwargs: Algorithm-specific parameters.
|
|
19
|
+
|
|
20
|
+
Returns:
|
|
21
|
+
A sampler instance.
|
|
22
|
+
|
|
23
|
+
Raises:
|
|
24
|
+
ValueError: If the algorithm name is unknown.
|
|
25
|
+
"""
|
|
26
|
+
kwargs = {k: v for k, v in kwargs.items() if v is not None}
|
|
27
|
+
|
|
28
|
+
samplers: dict[str, type[BaseSampler]] = {
|
|
29
|
+
"ucb1": UCB1Sampler,
|
|
30
|
+
"random": RandomSampler,
|
|
31
|
+
"greedy": GreedySampler,
|
|
32
|
+
"island": IslandSampler,
|
|
33
|
+
}
|
|
34
|
+
|
|
35
|
+
if algorithm not in samplers:
|
|
36
|
+
raise ValueError(
|
|
37
|
+
f"Unknown sampling algorithm: {algorithm}. "
|
|
38
|
+
f"Available: {list(samplers.keys())}"
|
|
39
|
+
)
|
|
40
|
+
|
|
41
|
+
return samplers[algorithm](**kwargs)
|
|
@@ -0,0 +1,18 @@
|
|
|
1
|
+
"""Greedy sampling strategy."""
|
|
2
|
+
|
|
3
|
+
from typing import List, TYPE_CHECKING
|
|
4
|
+
|
|
5
|
+
from .base import BaseSampler
|
|
6
|
+
|
|
7
|
+
if TYPE_CHECKING:
|
|
8
|
+
from ...utils.structures import Node
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
class GreedySampler(BaseSampler):
|
|
12
|
+
"""Return the highest-scoring nodes first."""
|
|
13
|
+
|
|
14
|
+
def sample(self, nodes: List["Node"], n: int) -> List["Node"]:
|
|
15
|
+
if not nodes:
|
|
16
|
+
return []
|
|
17
|
+
sorted_nodes = sorted(nodes, key=lambda x: x.score, reverse=True)
|
|
18
|
+
return sorted_nodes[:n]
|