Compare commits
7
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
73dde99e3a | ||
|
|
ecc8944006 | ||
|
|
dc15e2d04b | ||
|
|
ff1fc25d77 | ||
|
|
6b67326c3b | ||
|
|
d57f8778ac | ||
|
|
5b2d5b5f6e |
@@ -38,6 +38,22 @@
|
|||||||
## Plan (предыдущая задача — выполнена)
|
## Plan (предыдущая задача — выполнена)
|
||||||
1–7. Docker E2E пайплайна «install как человек → serve → визард → воркер → джоб» — выполнено, см. Progress ниже.
|
1–7. Docker E2E пайплайна «install как человек → serve → визард → воркер → джоб» — выполнено, см. Progress ниже.
|
||||||
|
|
||||||
|
## Progress (день: конкурентность в агенте — параллелизм через нашу архитектуру)
|
||||||
|
- ✅ По просьбе «реализовать параллелизм через нашу архитектуру»: `WORKER_CONCURRENCY` (поле `concurrency` в конфиге визарда, шаг Machine) — агент регистрируется один раз и ведёт N циклов claim→выполнение→upload под одним worker id. N шардов обрабатываются параллельно на одной машине, используя таски координатора как единицу параллелизма; SDK не менялся.
|
||||||
|
- ✅ Реализация: `Config.Concurrency` + env; `Daemon.RunForever` → register once → N goroutine-циклов (общий счётчик MaxTasks под мьютексом); визард: поле «Concurrent task loops» + конфиг; тест с маркер-скриптом доказывает 3 параллельных исполнения (race-тесты зелёные после мьютекса в fake).
|
||||||
|
- ✅ Измерено на релизном коде в Docker (один воркер, 8 шардов): concurrency=1 → 12s; concurrency=4 → 4s (**3×**). Плюс прежний `similarity-search-parallel` (потоки внутри шарда) композируется с конкурентностью.
|
||||||
|
- ✅ Релиз v1.1.0-alpha.20 (бинарники + wheel); полный гейт: race + lint 0 issues + pytest 213.
|
||||||
|
- ✅ Бинарник worker-agent на машине пользователя обновлён до alpha.20.
|
||||||
|
- GIL-высвобождения в RDKit нет (проверено: `RDKIT_ALLOW_GIL_RELEASE` не помогает), поэтому внутришардовые потоки не ускоряют чистый RDKit-путь — конкурентность задач это и компенсирует.
|
||||||
|
|
||||||
|
## Progress (утро: similarity-search-parallel)
|
||||||
|
- ✅ Новый workload `similarity-search-parallel@1.0.0` (отдельная версия, как просил пользователь): подкласс `SimilaritySearchSDKWorkload` + параллельное ядро `search_parallel/core.py` — fingerprinting+скоринг шарда через `ThreadPoolExecutor` (параметр `threads`, default CPU count); `pool.map` сохраняет порядок строк, поэтому merge идентичен последовательному (`_HeapEntry`) и результат **байт-в-байт** равен `similarity-search` при любом числе потоков.
|
||||||
|
- ✅ Тесты: байт-в-байт vs эталон для threads 1/2/4 с намеренными связями (изомеры, дубликаты), executor-прогон, валидация параметров, регистрация манифеста; 213 pytest зелёные.
|
||||||
|
- ✅ Экспортирован в каталог координатора (5 ворклоадов), Go-тесты зелёные; release v1.1.0-alpha.19 (бинарники + wheel `scimesh-1.1.0a19`).
|
||||||
|
- ✅ Распределённый E2E в Docker: воркер с wheel a19, джоб similarity-search-parallel (threads=4) → completed → результат байт-в-байт = локальному эталону.
|
||||||
|
- ✅ Бинарники на машине пользователя обновлены до alpha.19 (для появления ворклоада в UI нужен рестарт serve + переустановка рантайма воркера).
|
||||||
|
- Примечание: в CPython потоки не ускоряют чистый RDKit-путь (GIL), но структура готова к ядрам, отпускающим GIL (numpy и т.п.); при желании можно добавить процесс-пул как отдельный workload в будущей версии протокола.
|
||||||
|
|
||||||
## Progress (ночная сессия)
|
## Progress (ночная сессия)
|
||||||
- ✅ **П.1 Пустое имя воркера**: `domain.NewWorker` нормализует/отклоняет пустое имя + `TestNewWorkerRejectsBlankName` (починен `fixedTime`→`testNow`).
|
- ✅ **П.1 Пустое имя воркера**: `domain.NewWorker` нормализует/отклоняет пустое имя + `TestNewWorkerRejectsBlankName` (починен `fixedTime`→`testNow`).
|
||||||
- ✅ **П.2 `--check`**: пробует managed venv (если установлен) + реальный пробинг учётки (exchange ключа / claim-пробa) — `CheckAuth` + тесты; на машине пользователя: `✓ auth: credential accepted`, venv python, scimesh installed.
|
- ✅ **П.2 `--check`**: пробует managed venv (если установлен) + реальный пробинг учётки (exchange ключа / claim-пробa) — `CheckAuth` + тесты; на машине пользователя: `✓ auth: credential accepted`, venv python, scimesh installed.
|
||||||
|
|||||||
@@ -45,7 +45,7 @@ func runAgent(args []string) error {
|
|||||||
return fmt.Errorf("--coordinator-url, --token, and --work-dir are required")
|
return fmt.Errorf("--coordinator-url, --token, and --work-dir are required")
|
||||||
}
|
}
|
||||||
if *taskRunner == "" {
|
if *taskRunner == "" {
|
||||||
*taskRunner = "python -m scimesh.worker.task"
|
*taskRunner = "python -I -m scimesh.worker.task"
|
||||||
}
|
}
|
||||||
|
|
||||||
logger := slog.New(slog.NewTextHandler(os.Stderr, nil))
|
logger := slog.New(slog.NewTextHandler(os.Stderr, nil))
|
||||||
|
|||||||
@@ -236,9 +236,9 @@ func stopAgents(agents []*exec.Cmd) {
|
|||||||
// system `python`.
|
// system `python`.
|
||||||
func defaultTaskRunner(venvPython string) string {
|
func defaultTaskRunner(venvPython string) string {
|
||||||
if runtimeStatus(venvPython) {
|
if runtimeStatus(venvPython) {
|
||||||
return venvPython + " -m scimesh.worker.task"
|
return venvPython + " -I -m scimesh.worker.task"
|
||||||
}
|
}
|
||||||
return "python -m scimesh.worker.task"
|
return "python -I -m scimesh.worker.task"
|
||||||
}
|
}
|
||||||
|
|
||||||
// ensureRuntime creates the managed venv and installs scimesh into it, unless
|
// ensureRuntime creates the managed venv and installs scimesh into it, unless
|
||||||
|
|||||||
@@ -88,8 +88,12 @@ func CheckEnvironment(ctx context.Context) CheckReport {
|
|||||||
// execute with.
|
// execute with.
|
||||||
func CheckEnvironmentWithPython(ctx context.Context, python string) CheckReport {
|
func CheckEnvironmentWithPython(ctx context.Context, python string) CheckReport {
|
||||||
report := CheckReport{Agent: Version, Python: CheckItem{Name: "python", OK: true, Detail: python}}
|
report := CheckReport{Agent: Version, Python: CheckItem{Name: "python", OK: true, Detail: python}}
|
||||||
|
// The version comes from importlib.metadata, so the wizard can compare the
|
||||||
|
// installed package with the binary version and offer an upgrade. -I keeps
|
||||||
|
// the working directory out of sys.path, so a scimesh checkout in the
|
||||||
|
// wizard's cwd can never shadow the venv installation.
|
||||||
//nolint:gosec // G204: python is a resolved interpreter path, the argument list is constant
|
//nolint:gosec // G204: python is a resolved interpreter path, the argument list is constant
|
||||||
cmd := exec.CommandContext(ctx, python, "-c", "import scimesh; print(scimesh.__version__ if hasattr(scimesh, '__version__') else 'installed')")
|
cmd := exec.CommandContext(ctx, python, "-I", "-c", "import importlib.metadata as m; print(m.version('scimesh'))")
|
||||||
out, err := cmd.Output()
|
out, err := cmd.Output()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
// The worker executes workloads by spawning scimesh's task runner, so
|
// The worker executes workloads by spawning scimesh's task runner, so
|
||||||
|
|||||||
@@ -32,6 +32,10 @@ type Config struct {
|
|||||||
TaskRunner []string // command + args; defaults to python -m scimesh.worker.task
|
TaskRunner []string // command + args; defaults to python -m scimesh.worker.task
|
||||||
MaxTasks int // 0 = unlimited
|
MaxTasks int // 0 = unlimited
|
||||||
ExitWhenIdle bool
|
ExitWhenIdle bool
|
||||||
|
// Concurrency is how many claim→execute→upload loops run in parallel
|
||||||
|
// under one worker id: N shards processed concurrently on one machine,
|
||||||
|
// using the coordinator's own task pipeline as the parallel unit.
|
||||||
|
Concurrency int
|
||||||
}
|
}
|
||||||
|
|
||||||
func envList(name string) ([]string, error) {
|
func envList(name string) ([]string, error) {
|
||||||
@@ -108,7 +112,7 @@ func LoadConfig() (*Config, error) {
|
|||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
if len(runner) == 0 {
|
if len(runner) == 0 {
|
||||||
runner = []string{"python", "-m", "scimesh.worker.task"}
|
runner = []string{"python", "-I", "-m", "scimesh.worker.task"}
|
||||||
}
|
}
|
||||||
maxTasks := 0
|
maxTasks := 0
|
||||||
if raw := os.Getenv("MAX_TASKS"); raw != "" {
|
if raw := os.Getenv("MAX_TASKS"); raw != "" {
|
||||||
@@ -145,9 +149,22 @@ func LoadConfig() (*Config, error) {
|
|||||||
TaskRunner: runner,
|
TaskRunner: runner,
|
||||||
MaxTasks: maxTasks,
|
MaxTasks: maxTasks,
|
||||||
ExitWhenIdle: os.Getenv("EXIT_WHEN_IDLE") == "1",
|
ExitWhenIdle: os.Getenv("EXIT_WHEN_IDLE") == "1",
|
||||||
|
Concurrency: envInt("WORKER_CONCURRENCY", 1),
|
||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func envInt(name string, fallback int) int {
|
||||||
|
raw := os.Getenv(name)
|
||||||
|
if raw == "" {
|
||||||
|
return fallback
|
||||||
|
}
|
||||||
|
parsed, err := strconv.Atoi(raw)
|
||||||
|
if err != nil || parsed < 1 {
|
||||||
|
return fallback
|
||||||
|
}
|
||||||
|
return parsed
|
||||||
|
}
|
||||||
|
|
||||||
func durationEnv(name string, fallback time.Duration) (time.Duration, error) {
|
func durationEnv(name string, fallback time.Duration) (time.Duration, error) {
|
||||||
raw := os.Getenv(name)
|
raw := os.Getenv(name)
|
||||||
if raw == "" {
|
if raw == "" {
|
||||||
|
|||||||
@@ -21,6 +21,7 @@ type ConfigFile struct {
|
|||||||
WorkerName string `json:"worker_name,omitempty"`
|
WorkerName string `json:"worker_name,omitempty"`
|
||||||
CPUCount int `json:"cpu_count"`
|
CPUCount int `json:"cpu_count"`
|
||||||
MemoryMB int `json:"memory_mb"`
|
MemoryMB int `json:"memory_mb"`
|
||||||
|
Concurrency int `json:"concurrency,omitempty"`
|
||||||
TaskRunner []string `json:"task_runner,omitempty"`
|
TaskRunner []string `json:"task_runner,omitempty"`
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -103,11 +104,15 @@ func (f *ConfigFile) Config() (*Config, error) {
|
|||||||
if config.MemoryMB < 0 {
|
if config.MemoryMB < 0 {
|
||||||
config.MemoryMB = 0
|
config.MemoryMB = 0
|
||||||
}
|
}
|
||||||
|
config.Concurrency = f.Concurrency
|
||||||
|
if config.Concurrency < 1 {
|
||||||
|
config.Concurrency = 1
|
||||||
|
}
|
||||||
if len(f.TaskRunner) > 0 {
|
if len(f.TaskRunner) > 0 {
|
||||||
config.TaskRunner = f.TaskRunner
|
config.TaskRunner = f.TaskRunner
|
||||||
}
|
}
|
||||||
if len(config.TaskRunner) == 0 {
|
if len(config.TaskRunner) == 0 {
|
||||||
config.TaskRunner = []string{"python", "-m", "scimesh.worker.task"}
|
config.TaskRunner = []string{"python", "-I", "-m", "scimesh.worker.task"}
|
||||||
}
|
}
|
||||||
config.PollInterval = 2 * time.Second
|
config.PollInterval = 2 * time.Second
|
||||||
config.RequestTimeout = 30 * time.Second
|
config.RequestTimeout = 30 * time.Second
|
||||||
|
|||||||
@@ -37,15 +37,52 @@ func NewDaemon(config *Config, client *Client, runner *TaskRunner, log *slog.Log
|
|||||||
return &Daemon{config: config, client: client, runner: runner, log: log}
|
return &Daemon{config: config, client: client, runner: runner, log: log}
|
||||||
}
|
}
|
||||||
|
|
||||||
// RunForever loops until interrupted, idle-exit, or max tasks.
|
// RunForever registers once, then runs the claim→execute→upload loop
|
||||||
|
// concurrently under one worker id. With Concurrency > 1, several shards are
|
||||||
|
// processed in parallel on this machine, using the coordinator's own task
|
||||||
|
// pipeline as the parallel unit.
|
||||||
func (d *Daemon) RunForever() error {
|
func (d *Daemon) RunForever() error {
|
||||||
|
if !d.registered {
|
||||||
|
if err := d.register(); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
workers := d.config.Concurrency
|
||||||
|
if workers < 1 {
|
||||||
|
workers = 1
|
||||||
|
}
|
||||||
|
if workers == 1 {
|
||||||
|
return d.loop()
|
||||||
|
}
|
||||||
|
d.log.Info("agent running concurrently", "loops", workers)
|
||||||
|
var wg sync.WaitGroup
|
||||||
|
errs := make(chan error, workers)
|
||||||
|
for i := 0; i < workers; i++ {
|
||||||
|
wg.Add(1)
|
||||||
|
go func(loop int) {
|
||||||
|
defer wg.Done()
|
||||||
|
if err := d.loop(); err != nil {
|
||||||
|
errs <- err
|
||||||
|
return
|
||||||
|
}
|
||||||
|
errs <- nil
|
||||||
|
}(i)
|
||||||
|
}
|
||||||
|
wg.Wait()
|
||||||
|
close(errs)
|
||||||
|
for err := range errs {
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// loop is one claim→execute→upload cycle until interrupted, idle-exit, or
|
||||||
|
// the shared max-tasks budget is consumed.
|
||||||
|
func (d *Daemon) loop() error {
|
||||||
failures := 0
|
failures := 0
|
||||||
for {
|
for {
|
||||||
if !d.registered {
|
|
||||||
if err := d.register(); err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
}
|
|
||||||
d.cleanupExpiredDirectories()
|
d.cleanupExpiredDirectories()
|
||||||
outcome, err := d.runOnce()
|
outcome, err := d.runOnce()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -63,8 +100,11 @@ func (d *Daemon) RunForever() error {
|
|||||||
}
|
}
|
||||||
failures = 0
|
failures = 0
|
||||||
if outcome.Claimed && outcome.Completed {
|
if outcome.Claimed && outcome.Completed {
|
||||||
|
d.mu.Lock()
|
||||||
d.completed++
|
d.completed++
|
||||||
if d.config.MaxTasks > 0 && d.completed >= d.config.MaxTasks {
|
done := d.config.MaxTasks > 0 && d.completed >= d.config.MaxTasks
|
||||||
|
d.mu.Unlock()
|
||||||
|
if done {
|
||||||
d.log.Info("max tasks reached")
|
d.log.Info("max tasks reached")
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -11,6 +11,7 @@ import (
|
|||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
"strings"
|
"strings"
|
||||||
|
"sync"
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
)
|
)
|
||||||
@@ -55,6 +56,7 @@ type fakeCoordinator struct {
|
|||||||
uploadSize int64
|
uploadSize int64
|
||||||
inputBytes []byte
|
inputBytes []byte
|
||||||
conflict bool // 409 on heartbeat/upload/result
|
conflict bool // 409 on heartbeat/upload/result
|
||||||
|
mu sync.Mutex
|
||||||
}
|
}
|
||||||
|
|
||||||
func newFakeCoordinator(t *testing.T, task map[string]any) *fakeCoordinator {
|
func newFakeCoordinator(t *testing.T, task map[string]any) *fakeCoordinator {
|
||||||
@@ -64,6 +66,8 @@ func newFakeCoordinator(t *testing.T, task map[string]any) *fakeCoordinator {
|
|||||||
fake.uploadSize = int64(len(fake.inputBytes))
|
fake.uploadSize = int64(len(fake.inputBytes))
|
||||||
var server *httptest.Server
|
var server *httptest.Server
|
||||||
server = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
server = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
fake.mu.Lock()
|
||||||
|
defer fake.mu.Unlock()
|
||||||
switch {
|
switch {
|
||||||
case r.Method == http.MethodPost && r.URL.Path == "/workers/register":
|
case r.Method == http.MethodPost && r.URL.Path == "/workers/register":
|
||||||
writeJSON(w, http.StatusCreated, map[string]any{
|
writeJSON(w, http.StatusCreated, map[string]any{
|
||||||
@@ -294,3 +298,78 @@ func TestDaemonIdleClaimIsNotCompleted(t *testing.T) {
|
|||||||
t.Fatalf("outcome = %+v", outcome)
|
t.Fatalf("outcome = %+v", outcome)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestDaemonConcurrencyProcessesTasksInParallel(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
marker := filepath.Join(t.TempDir(), "marker")
|
||||||
|
script := filepath.Join(t.TempDir(), "fake-runner.sh")
|
||||||
|
content := `#!/bin/sh
|
||||||
|
out=""
|
||||||
|
task_dir=""
|
||||||
|
while [ "$#" -gt 0 ]; do
|
||||||
|
case "$1" in
|
||||||
|
--output) out="$2"; shift 2;;
|
||||||
|
--task-dir) task_dir="$2"; shift 2;;
|
||||||
|
*) shift;;
|
||||||
|
esac
|
||||||
|
done
|
||||||
|
echo start >> ` + marker + `
|
||||||
|
sleep 1
|
||||||
|
echo end >> ` + marker + `
|
||||||
|
printf 'id,score\n1,1\n' > "$task_dir/result.csv"
|
||||||
|
printf '{"artifact_path":"%s/result.csv","content_type":"text/csv","metrics":{"rows":1}}' "$task_dir" > "$out"
|
||||||
|
exit 0
|
||||||
|
`
|
||||||
|
if err := os.WriteFile(script, []byte(content), 0o755); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
fake := newFakeCoordinator(t, validClaimedTaskPayload())
|
||||||
|
defer fake.close()
|
||||||
|
config := &Config{
|
||||||
|
CoordinatorURL: fake.server.URL,
|
||||||
|
WorkerName: "concurrent-worker",
|
||||||
|
WorkerID: "22222222-2222-4222-8222-222222222222",
|
||||||
|
WorkDir: t.TempDir(),
|
||||||
|
CPUCount: 1,
|
||||||
|
PollInterval: time.Millisecond,
|
||||||
|
RequestTimeout: 5 * time.Second,
|
||||||
|
Heartbeat: 15 * time.Second,
|
||||||
|
Capabilities: []string{"similarity-search"},
|
||||||
|
TaskRunner: []string{script},
|
||||||
|
MaxTasks: 3,
|
||||||
|
Concurrency: 3,
|
||||||
|
}
|
||||||
|
client := NewClient(fake.server.URL, &StaticToken{token: "test-token"}, 5*time.Second)
|
||||||
|
logger := slog.New(slog.NewTextHandler(io.Discard, nil))
|
||||||
|
daemon := NewDaemon(config, client, NewTaskRunner(config.TaskRunner), logger)
|
||||||
|
if err := daemon.RunForever(); err != nil {
|
||||||
|
t.Fatalf("run: %v", err)
|
||||||
|
}
|
||||||
|
raw, err := os.ReadFile(marker)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("marker: %v", err)
|
||||||
|
}
|
||||||
|
starts := strings.Count(string(raw), "start\n")
|
||||||
|
ends := strings.Count(string(raw), "end\n")
|
||||||
|
if starts < 3 || ends < 3 {
|
||||||
|
t.Fatalf("marker: %d starts / %d ends, want at least 3/3", starts, ends)
|
||||||
|
}
|
||||||
|
// With three loops sleeping 1s each, the marker proves all three ran
|
||||||
|
// concurrently (three starts before the first end completes a 1s sleep).
|
||||||
|
lines := strings.Split(strings.TrimSpace(string(raw)), "\n")
|
||||||
|
concurrent := 0
|
||||||
|
running := 0
|
||||||
|
for _, line := range lines {
|
||||||
|
if line == "start" {
|
||||||
|
running++
|
||||||
|
if running > concurrent {
|
||||||
|
concurrent = running
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
running--
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if concurrent < 3 {
|
||||||
|
t.Errorf("max concurrent executions = %d, want 3", concurrent)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -323,7 +323,7 @@ func (s *Server) ensureVenvTaskRunner() {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
if venv := s.venvPython(); venv != "" {
|
if venv := s.venvPython(); venv != "" {
|
||||||
file.TaskRunner = []string{venv, "-m", "scimesh.worker.task"}
|
file.TaskRunner = []string{venv, "-I", "-m", "scimesh.worker.task"}
|
||||||
if payload, err := json.MarshalIndent(file, "", " "); err == nil {
|
if payload, err := json.MarshalIndent(file, "", " "); err == nil {
|
||||||
_ = os.WriteFile(s.cfgPath, append(payload, '\n'), 0o600)
|
_ = os.WriteFile(s.cfgPath, append(payload, '\n'), 0o600)
|
||||||
}
|
}
|
||||||
@@ -377,6 +377,7 @@ type saveConfigRequest struct {
|
|||||||
WorkerName string `json:"worker_name"`
|
WorkerName string `json:"worker_name"`
|
||||||
CPUCount int `json:"cpu_count"`
|
CPUCount int `json:"cpu_count"`
|
||||||
MemoryMB int `json:"memory_mb"`
|
MemoryMB int `json:"memory_mb"`
|
||||||
|
Concurrency int `json:"concurrency"`
|
||||||
TaskRunner []string `json:"task_runner"`
|
TaskRunner []string `json:"task_runner"`
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -395,6 +396,7 @@ func (s *Server) handleSaveConfig(w http.ResponseWriter, r *http.Request) {
|
|||||||
WorkerName: strings.TrimSpace(req.WorkerName),
|
WorkerName: strings.TrimSpace(req.WorkerName),
|
||||||
CPUCount: req.CPUCount,
|
CPUCount: req.CPUCount,
|
||||||
MemoryMB: req.MemoryMB,
|
MemoryMB: req.MemoryMB,
|
||||||
|
Concurrency: req.Concurrency,
|
||||||
TaskRunner: req.TaskRunner,
|
TaskRunner: req.TaskRunner,
|
||||||
}
|
}
|
||||||
if file.CoordinatorURL == "" {
|
if file.CoordinatorURL == "" {
|
||||||
@@ -418,12 +420,15 @@ func (s *Server) handleSaveConfig(w http.ResponseWriter, r *http.Request) {
|
|||||||
if file.CPUCount < 1 {
|
if file.CPUCount < 1 {
|
||||||
file.CPUCount = 1
|
file.CPUCount = 1
|
||||||
}
|
}
|
||||||
|
if file.Concurrency < 1 {
|
||||||
|
file.Concurrency = 1
|
||||||
|
}
|
||||||
// The wizard UI bakes the venv python into the runner after an install;
|
// The wizard UI bakes the venv python into the runner after an install;
|
||||||
// an API-driven or scripted flow may not, so the server guarantees it:
|
// an API-driven or scripted flow may not, so the server guarantees it:
|
||||||
// workloads execute through scimesh's task runner, which lives in the venv.
|
// workloads execute through scimesh's task runner, which lives in the venv.
|
||||||
if len(file.TaskRunner) == 0 {
|
if len(file.TaskRunner) == 0 {
|
||||||
if venv := s.venvPython(); venv != "" {
|
if venv := s.venvPython(); venv != "" {
|
||||||
file.TaskRunner = []string{venv, "-m", "scimesh.worker.task"}
|
file.TaskRunner = []string{venv, "-I", "-m", "scimesh.worker.task"}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
if err := agent.SaveConfigFile(s.cfgPath, file); err != nil {
|
if err := agent.SaveConfigFile(s.cfgPath, file); err != nil {
|
||||||
@@ -449,6 +454,9 @@ func (s *Server) handleTest(w http.ResponseWriter, r *http.Request) {
|
|||||||
// checking the bare system python3 would keep reporting scimesh as
|
// checking the bare system python3 would keep reporting scimesh as
|
||||||
// missing even though the worker would run with the venv.
|
// missing even though the worker would run with the venv.
|
||||||
report := agent.RunCheck(r.Context(), url, s.venvPython(), req.Token, req.WorkerKey, req.UserserviceURL)
|
report := agent.RunCheck(r.Context(), url, s.venvPython(), req.Token, req.WorkerKey, req.UserserviceURL)
|
||||||
|
if report.Scimesh.OK {
|
||||||
|
report.Scimesh = ensureMatchingScimeshVersion(report.Scimesh)
|
||||||
|
}
|
||||||
writeJSON(w, http.StatusOK, report)
|
writeJSON(w, http.StatusOK, report)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -623,3 +631,26 @@ func truncate(s string, n int) string {
|
|||||||
}
|
}
|
||||||
return s[:n] + "…"
|
return s[:n] + "…"
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// ensureMatchingScimeshVersion flips a green scimesh check to a stale one when
|
||||||
|
// the installed package does not match the worker-agent's own version: a
|
||||||
|
// version-locked wheel is the only supported runtime, and a mismatch means the
|
||||||
|
// workload catalog the worker advertises is not what it executes. The wizard
|
||||||
|
// UI then offers the Install button again. Dev builds have no release wheel,
|
||||||
|
// so they skip the comparison.
|
||||||
|
func ensureMatchingScimeshVersion(item agent.CheckItem) agent.CheckItem {
|
||||||
|
if agent.Version == "" || agent.Version == "dev" {
|
||||||
|
return item
|
||||||
|
}
|
||||||
|
want := agent.NormalizePEP440(agent.Version)
|
||||||
|
got := strings.TrimSpace(item.Detail)
|
||||||
|
if got == "" || got == want {
|
||||||
|
return item
|
||||||
|
}
|
||||||
|
item.OK = false
|
||||||
|
item.Detail = fmt.Sprintf(
|
||||||
|
"installed scimesh %s, but this worker-agent (%s) needs %s — press Install to upgrade",
|
||||||
|
got, agent.Version, want,
|
||||||
|
)
|
||||||
|
return item
|
||||||
|
}
|
||||||
|
|||||||
@@ -429,8 +429,8 @@ func TestStartPinsTheVenvTaskRunner(t *testing.T) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
if len(config.TaskRunner) != 3 || config.TaskRunner[0] != venvPython || config.TaskRunner[1] != "-m" || config.TaskRunner[2] != "scimesh.worker.task" {
|
if len(config.TaskRunner) != 4 || config.TaskRunner[0] != venvPython || config.TaskRunner[1] != "-I" || config.TaskRunner[2] != "-m" || config.TaskRunner[3] != "scimesh.worker.task" {
|
||||||
t.Errorf("task runner = %v, want the venv python runner", config.TaskRunner)
|
t.Errorf("task runner = %v, want the venv python runner with -I", config.TaskRunner)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -450,8 +450,8 @@ func TestSaveConfigPinsVenvRunnerWhenPresent(t *testing.T) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
if len(config.TaskRunner) != 3 || config.TaskRunner[0] != venvPython {
|
if len(config.TaskRunner) != 4 || config.TaskRunner[0] != venvPython || config.TaskRunner[1] != "-I" {
|
||||||
t.Errorf("task runner = %v, want the venv python", config.TaskRunner)
|
t.Errorf("task runner = %v, want the venv python with -I", config.TaskRunner)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -462,7 +462,7 @@ func TestTestProbesTheVenvPythonAfterInstall(t *testing.T) {
|
|||||||
// a fake scimesh version so the preflight goes green through the venv.
|
// a fake scimesh version so the preflight goes green through the venv.
|
||||||
venvPython := filepath.Join(server.dir, "venv", "bin", "python")
|
venvPython := filepath.Join(server.dir, "venv", "bin", "python")
|
||||||
_ = os.MkdirAll(filepath.Dir(venvPython), 0o755)
|
_ = os.MkdirAll(filepath.Dir(venvPython), 0o755)
|
||||||
_ = os.WriteFile(venvPython, []byte("#!/bin/sh\nif [ \"$1\" = \"-c\" ]; then echo 9.9.9-test; exit 0; fi\nexit 0\n"), 0o755)
|
_ = os.WriteFile(venvPython, []byte("#!/bin/sh\nfor a in \"$@\"; do if [ \"$a\" = \"-c\" ]; then echo 9.9.9-test; exit 0; fi; done\nexit 0\n"), 0o755)
|
||||||
|
|
||||||
req, _ := http.NewRequestWithContext(context.Background(), http.MethodPost, base+"/api/test", strings.NewReader(`{"coordinator_url":"http://127.0.0.1:1"}`))
|
req, _ := http.NewRequestWithContext(context.Background(), http.MethodPost, base+"/api/test", strings.NewReader(`{"coordinator_url":"http://127.0.0.1:1"}`))
|
||||||
req.Header.Set("Content-Type", "application/json")
|
req.Header.Set("Content-Type", "application/json")
|
||||||
@@ -497,3 +497,26 @@ time=6 level=WARN msg="agent cycle failed" error="boom"
|
|||||||
t.Errorf("stats = %+v, want registered claimed=2 completed=1 failed=1", stats)
|
t.Errorf("stats = %+v, want registered claimed=2 completed=1 failed=1", stats)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func testCheckScimeshVersion(t *testing.T, installed, binary string, wantOK bool, wantDetail string) {
|
||||||
|
t.Helper()
|
||||||
|
old := agent.Version
|
||||||
|
agent.Version = binary
|
||||||
|
t.Cleanup(func() { agent.Version = old })
|
||||||
|
item := ensureMatchingScimeshVersion(agent.CheckItem{Name: "scimesh", OK: true, Detail: installed})
|
||||||
|
if item.OK != wantOK {
|
||||||
|
t.Errorf("installed=%s binary=%s: ok=%v, want %v (%s)", installed, binary, item.OK, wantOK, item.Detail)
|
||||||
|
}
|
||||||
|
if wantDetail != "" && !strings.Contains(item.Detail, wantDetail) {
|
||||||
|
t.Errorf("detail = %q, want it to contain %q", item.Detail, wantDetail)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestEnsureMatchingScimeshVersion(t *testing.T) {
|
||||||
|
testCheckScimeshVersion(t, "1.1.0a20", "1.1.0-alpha.20", true, "")
|
||||||
|
testCheckScimeshVersion(t, "1.1.0a17", "1.1.0-alpha.20", false, "press Install to upgrade")
|
||||||
|
testCheckScimeshVersion(t, "1.1.0a16.dev7+gea0fb8c59.d20260803", "1.1.0-alpha.20", false, "needs 1.1.0a20")
|
||||||
|
// Dev builds and unknown versions never block.
|
||||||
|
testCheckScimeshVersion(t, "anything", "dev", true, "")
|
||||||
|
testCheckScimeshVersion(t, "1.1.0a20", "", true, "")
|
||||||
|
}
|
||||||
|
|||||||
@@ -141,6 +141,7 @@ code{font-family:var(--mono);font-size:.86em}
|
|||||||
</div>
|
</div>
|
||||||
</div>
|
</div>
|
||||||
<div class="field" id="cpu-field" style="display:none"><label>CPU count</label><input id="in-cpu" type="number" min="1" value="1"></div>
|
<div class="field" id="cpu-field" style="display:none"><label>CPU count</label><input id="in-cpu" type="number" min="1" value="1"></div>
|
||||||
|
<div class="field"><label>Concurrent task loops</label><input id="in-conc" type="number" min="1" max="64" value="1"><p class="hint">Process this many shards in parallel on this machine. Each loop runs its own task runner subprocess.</p></div>
|
||||||
<div class="actions"><button class="btn btn-ghost" id="b2b">← Back</button><button class="btn btn-primary" id="b2">Continue →</button></div>
|
<div class="actions"><button class="btn btn-ghost" id="b2b">← Back</button><button class="btn btn-primary" id="b2">Continue →</button></div>
|
||||||
</div>
|
</div>
|
||||||
</div>
|
</div>
|
||||||
@@ -228,9 +229,10 @@ function draftConfig(){
|
|||||||
userservice_url:state.mode==='key'?$('in-users').value.trim():'',
|
userservice_url:state.mode==='key'?$('in-users').value.trim():'',
|
||||||
work_dir:$('in-dir').value.trim(),
|
work_dir:$('in-dir').value.trim(),
|
||||||
worker_name:$('in-name').value.trim(),
|
worker_name:$('in-name').value.trim(),
|
||||||
cpu_count:state.cpu==='custom'?parseInt($('in-cpu').value||'1',10):0
|
cpu_count:state.cpu==='custom'?parseInt($('in-cpu').value||'1',10):0,
|
||||||
|
concurrency:parseInt($('in-conc').value||'1',10)
|
||||||
};
|
};
|
||||||
if(state.venvPython)cfg.task_runner=[state.venvPython,'-m','scimesh.worker.task'];
|
if(state.venvPython)cfg.task_runner=[state.venvPython,'-I','-m','scimesh.worker.task'];
|
||||||
return cfg;
|
return cfg;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -597,6 +597,206 @@
|
|||||||
"upload_ready": true,
|
"upload_ready": true,
|
||||||
"verifier": "exact-artifact@1",
|
"verifier": "exact-artifact@1",
|
||||||
"version": "1.0.0"
|
"version": "1.0.0"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"capabilities": [
|
||||||
|
"similarity-search-parallel"
|
||||||
|
],
|
||||||
|
"description": "Exact top-k Tanimoto molecular similarity search over deterministic TSV shards with a bounded merge; each shard is fingerprinted and scored across a thread pool. Output is byte-identical to similarity-search.",
|
||||||
|
"determinism": "byte_exact",
|
||||||
|
"enabled": true,
|
||||||
|
"inputs": {
|
||||||
|
"input": {
|
||||||
|
"allow_nested_collections": false,
|
||||||
|
"canonicalizer": "scimesh-tsv-v1",
|
||||||
|
"encoding": "utf-8",
|
||||||
|
"max_bytes": 10737418240,
|
||||||
|
"max_dimensions": [],
|
||||||
|
"max_records": 100000000,
|
||||||
|
"media_type": "text/tab-separated-values",
|
||||||
|
"privacy_class": "project",
|
||||||
|
"ref": "molecule-table@1",
|
||||||
|
"retention_class": "durable",
|
||||||
|
"streaming": false,
|
||||||
|
"validator": "delimited-table@1",
|
||||||
|
"validator_configuration": {
|
||||||
|
"required_columns": [
|
||||||
|
"canonical_smiles",
|
||||||
|
"chembl_id"
|
||||||
|
]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"name": "similarity-search-parallel",
|
||||||
|
"outputs": {
|
||||||
|
"result": {
|
||||||
|
"allow_nested_collections": false,
|
||||||
|
"canonicalizer": "scimesh-search-result-v1",
|
||||||
|
"encoding": "utf-8",
|
||||||
|
"max_bytes": 1073741824,
|
||||||
|
"max_dimensions": [],
|
||||||
|
"max_records": 100000,
|
||||||
|
"media_type": "text/csv",
|
||||||
|
"privacy_class": "project",
|
||||||
|
"ref": "similarity-search-result@1",
|
||||||
|
"retention_class": "durable",
|
||||||
|
"streaming": false,
|
||||||
|
"validator": "delimited-table@1",
|
||||||
|
"validator_configuration": {
|
||||||
|
"columns": [
|
||||||
|
"rank",
|
||||||
|
"chembl_id",
|
||||||
|
"canonical_smiles",
|
||||||
|
"similarity"
|
||||||
|
]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"parameters_schema": {
|
||||||
|
"additionalProperties": false,
|
||||||
|
"oneOf": [
|
||||||
|
{
|
||||||
|
"not": {
|
||||||
|
"required": [
|
||||||
|
"query_smiles"
|
||||||
|
]
|
||||||
|
},
|
||||||
|
"required": [
|
||||||
|
"query_id"
|
||||||
|
]
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"not": {
|
||||||
|
"required": [
|
||||||
|
"query_id"
|
||||||
|
]
|
||||||
|
},
|
||||||
|
"required": [
|
||||||
|
"query_smiles"
|
||||||
|
]
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"properties": {
|
||||||
|
"max_rows": {
|
||||||
|
"minimum": 1,
|
||||||
|
"type": "integer"
|
||||||
|
},
|
||||||
|
"progress_every": {
|
||||||
|
"minimum": 0,
|
||||||
|
"type": "integer"
|
||||||
|
},
|
||||||
|
"query_id": {
|
||||||
|
"maxLength": 200,
|
||||||
|
"minLength": 1,
|
||||||
|
"type": "string"
|
||||||
|
},
|
||||||
|
"query_smiles": {
|
||||||
|
"maxLength": 200,
|
||||||
|
"minLength": 1,
|
||||||
|
"type": "string"
|
||||||
|
},
|
||||||
|
"threads": {
|
||||||
|
"description": "Threads used to fingerprint and score one shard (default: CPU count).",
|
||||||
|
"minimum": 1,
|
||||||
|
"type": "integer"
|
||||||
|
},
|
||||||
|
"threshold": {
|
||||||
|
"maximum": 1,
|
||||||
|
"minimum": 0,
|
||||||
|
"type": "number"
|
||||||
|
},
|
||||||
|
"threshold_direction": {
|
||||||
|
"enum": [
|
||||||
|
"greater",
|
||||||
|
"less"
|
||||||
|
]
|
||||||
|
},
|
||||||
|
"top_k": {
|
||||||
|
"minimum": 1,
|
||||||
|
"type": "integer"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"type": "object"
|
||||||
|
},
|
||||||
|
"reduction": "top-k",
|
||||||
|
"trust_modes": [
|
||||||
|
"trusted",
|
||||||
|
"untrusted_quorum"
|
||||||
|
],
|
||||||
|
"ui_elements": [
|
||||||
|
{
|
||||||
|
"default": null,
|
||||||
|
"field": "query_id",
|
||||||
|
"group": "",
|
||||||
|
"help": "ChEMBL id of the query molecule. Provide exactly one of id or SMILES.",
|
||||||
|
"label": "Query molecule id",
|
||||||
|
"options": [],
|
||||||
|
"order": 1,
|
||||||
|
"placeholder": "",
|
||||||
|
"widget": "text"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"default": null,
|
||||||
|
"field": "query_smiles",
|
||||||
|
"group": "",
|
||||||
|
"help": "SMILES of the query molecule. Provide exactly one of id or SMILES.",
|
||||||
|
"label": "Query molecule SMILES",
|
||||||
|
"options": [],
|
||||||
|
"order": 2,
|
||||||
|
"placeholder": "",
|
||||||
|
"widget": "text"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"default": 20,
|
||||||
|
"field": "top_k",
|
||||||
|
"group": "",
|
||||||
|
"help": "Number of most similar molecules to keep per shard (global merge keeps the best of these).",
|
||||||
|
"label": "Top k",
|
||||||
|
"options": [],
|
||||||
|
"order": 3,
|
||||||
|
"placeholder": "",
|
||||||
|
"widget": "number"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"default": "greater",
|
||||||
|
"field": "threshold_direction",
|
||||||
|
"group": "",
|
||||||
|
"help": "Keep molecules with similarity greater or less than the threshold.",
|
||||||
|
"label": "Direction",
|
||||||
|
"options": [
|
||||||
|
"greater",
|
||||||
|
"less"
|
||||||
|
],
|
||||||
|
"order": 4,
|
||||||
|
"placeholder": "",
|
||||||
|
"widget": "select"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"default": null,
|
||||||
|
"field": "threshold",
|
||||||
|
"group": "",
|
||||||
|
"help": "Optional similarity bound: results are filtered to this direction.",
|
||||||
|
"label": "Similarity threshold",
|
||||||
|
"options": [],
|
||||||
|
"order": 5,
|
||||||
|
"placeholder": "e.g. 0.8",
|
||||||
|
"widget": "number"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"default": null,
|
||||||
|
"field": "threads",
|
||||||
|
"group": "",
|
||||||
|
"help": "Threads used to fingerprint and score one shard (default: CPU count).",
|
||||||
|
"label": "Threads per shard",
|
||||||
|
"options": [],
|
||||||
|
"order": 6,
|
||||||
|
"placeholder": "auto",
|
||||||
|
"widget": "number"
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"upload_ready": true,
|
||||||
|
"verifier": "exact-artifact@1",
|
||||||
|
"version": "1.0.0"
|
||||||
}
|
}
|
||||||
]
|
]
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -23,6 +23,7 @@ scimesh = "scimesh.cli:main"
|
|||||||
|
|
||||||
[project.entry-points."scimesh.workloads"]
|
[project.entry-points."scimesh.workloads"]
|
||||||
"similarity-search@1.0.0" = "scimesh.workloads.search:workload_definition"
|
"similarity-search@1.0.0" = "scimesh.workloads.search:workload_definition"
|
||||||
|
"similarity-search-parallel@1.0.0" = "scimesh.workloads.search_parallel:workload_definition"
|
||||||
"similarity-graph@1.0.0" = "scimesh.workloads.graph:workload_definition"
|
"similarity-graph@1.0.0" = "scimesh.workloads.graph:workload_definition"
|
||||||
"descriptor-batch@1.0.0" = "scimesh.workloads.descriptors:workload_definition"
|
"descriptor-batch@1.0.0" = "scimesh.workloads.descriptors:workload_definition"
|
||||||
"molwt-filter@1.0.0" = "scimesh.workloads.molwt_filter:workload_definition"
|
"molwt-filter@1.0.0" = "scimesh.workloads.molwt_filter:workload_definition"
|
||||||
|
|||||||
@@ -21,6 +21,7 @@ from .environment import current_environment_digest
|
|||||||
from .graph import similarity_graph_sdk_definition
|
from .graph import similarity_graph_sdk_definition
|
||||||
from .molwt_filter import molwt_filter_sdk_definition
|
from .molwt_filter import molwt_filter_sdk_definition
|
||||||
from .search import similarity_search_sdk_definition
|
from .search import similarity_search_sdk_definition
|
||||||
|
from .search_parallel import similarity_search_parallel_sdk_definition
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
"default_sdk_registry",
|
"default_sdk_registry",
|
||||||
@@ -46,6 +47,10 @@ def default_sdk_registry(
|
|||||||
similarity_search_sdk_definition(shard_rows=shard_rows).definition(),
|
similarity_search_sdk_definition(shard_rows=shard_rows).definition(),
|
||||||
enabled=True,
|
enabled=True,
|
||||||
)
|
)
|
||||||
|
registry.register(
|
||||||
|
similarity_search_parallel_sdk_definition(shard_rows=shard_rows).definition(),
|
||||||
|
enabled=True,
|
||||||
|
)
|
||||||
registry.register(
|
registry.register(
|
||||||
similarity_graph_sdk_definition().definition(),
|
similarity_graph_sdk_definition().definition(),
|
||||||
enabled=True,
|
enabled=True,
|
||||||
@@ -82,6 +87,7 @@ def default_sdk_runtime(
|
|||||||
workload_capabilities
|
workload_capabilities
|
||||||
or (
|
or (
|
||||||
"similarity-search",
|
"similarity-search",
|
||||||
|
"similarity-search-parallel",
|
||||||
"similarity-graph",
|
"similarity-graph",
|
||||||
"descriptor-batch",
|
"descriptor-batch",
|
||||||
"molwt-filter",
|
"molwt-filter",
|
||||||
|
|||||||
@@ -0,0 +1,28 @@
|
|||||||
|
"""SDK-built ``similarity-search-parallel`` workload.
|
||||||
|
|
||||||
|
Same contract as ``similarity-search`` with a per-shard thread pool. See
|
||||||
|
``core.py`` for the parallel scoring core and ``definition.py`` for the
|
||||||
|
manifest-backed handlers.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from .core import (
|
||||||
|
run_search_shard_parallel,
|
||||||
|
search_similar_parallel,
|
||||||
|
write_search_shards,
|
||||||
|
)
|
||||||
|
from .definition import (
|
||||||
|
MAP_ENTRY_POINT,
|
||||||
|
SimilaritySearchParallelSDKWorkload,
|
||||||
|
similarity_search_parallel_sdk_definition,
|
||||||
|
workload_definition,
|
||||||
|
)
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"MAP_ENTRY_POINT",
|
||||||
|
"SimilaritySearchParallelSDKWorkload",
|
||||||
|
"similarity_search_parallel_sdk_definition",
|
||||||
|
"workload_definition",
|
||||||
|
"run_search_shard_parallel",
|
||||||
|
"search_similar_parallel",
|
||||||
|
"write_search_shards",
|
||||||
|
]
|
||||||
@@ -0,0 +1,199 @@
|
|||||||
|
"""Scientific core for the SDK-built ``similarity-search-parallel`` workload.
|
||||||
|
|
||||||
|
The exact same semantics as ``similarity-search`` — identical partial format,
|
||||||
|
identical bounded merge, byte-identical output — but the per-molecule
|
||||||
|
fingerprinting and Tanimoto scoring of one shard run across a thread pool
|
||||||
|
(``threads`` parameter, default = CPU count).
|
||||||
|
|
||||||
|
Parallelism is confined to the scoring phase: ``ThreadPoolExecutor.map`` keeps
|
||||||
|
the input row order, so the results are merged exactly like the sequential
|
||||||
|
reference (same ``_HeapEntry`` logic), which makes the output byte-identical
|
||||||
|
for every thread count by construction.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import heapq
|
||||||
|
import os
|
||||||
|
from concurrent.futures import ThreadPoolExecutor
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Mapping
|
||||||
|
|
||||||
|
from rdkit import Chem
|
||||||
|
from rdkit.Chem import DataStructs
|
||||||
|
|
||||||
|
from scimesh.chemistry.dataset import MoleculeRecord, parse_smiles
|
||||||
|
from scimesh.chemistry.fingerprints import fingerprint
|
||||||
|
from scimesh.workloads.search.core import (
|
||||||
|
run_search_shard,
|
||||||
|
write_search_partial,
|
||||||
|
write_search_shards,
|
||||||
|
)
|
||||||
|
from scimesh.workloads.similarity_search import (
|
||||||
|
DatasetStats,
|
||||||
|
SearchResult,
|
||||||
|
SimilarityMatch,
|
||||||
|
_HeapEntry,
|
||||||
|
iter_valid_molecules,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def search_similar_parallel(
|
||||||
|
tsv_path: Path,
|
||||||
|
query: MoleculeRecord,
|
||||||
|
top_k: int,
|
||||||
|
*,
|
||||||
|
threads: int = 0,
|
||||||
|
max_rows: int | None = None,
|
||||||
|
threshold: float | None = None,
|
||||||
|
threshold_direction: str = "greater",
|
||||||
|
) -> SearchResult:
|
||||||
|
"""Exact top-k matches with a bounded heap, scored by a thread pool.
|
||||||
|
|
||||||
|
Identical selection and ordering to ``search_similar`` for every thread
|
||||||
|
count: the merge runs in row order over the parallel-computed scores.
|
||||||
|
"""
|
||||||
|
if top_k < 1:
|
||||||
|
raise ValueError("--top-k must be a positive integer")
|
||||||
|
if threads < 0:
|
||||||
|
raise ValueError("threads must be a non-negative integer")
|
||||||
|
if threshold is not None and not 0.0 <= threshold <= 1.0:
|
||||||
|
raise ValueError("--threshold must be between 0 and 1")
|
||||||
|
if threshold_direction not in {"greater", "less"}:
|
||||||
|
raise ValueError("--threshold-direction must be 'greater' or 'less'")
|
||||||
|
workers = threads or (os.cpu_count() or 1)
|
||||||
|
|
||||||
|
query_fingerprint = fingerprint(query.molecule)
|
||||||
|
query_canonical_smiles = Chem.MolToSmiles(query.molecule, canonical=True)
|
||||||
|
stats = DatasetStats()
|
||||||
|
records = list(iter_valid_molecules(tsv_path, stats, max_rows=max_rows))
|
||||||
|
|
||||||
|
def score(record: MoleculeRecord):
|
||||||
|
candidate_smiles = Chem.MolToSmiles(record.molecule, canonical=True)
|
||||||
|
if (
|
||||||
|
record.molecule_id == query.molecule_id
|
||||||
|
or candidate_smiles == query_canonical_smiles
|
||||||
|
):
|
||||||
|
return None
|
||||||
|
similarity = DataStructs.TanimotoSimilarity(
|
||||||
|
query_fingerprint, fingerprint(record.molecule)
|
||||||
|
)
|
||||||
|
if threshold is not None and (
|
||||||
|
similarity < threshold
|
||||||
|
if threshold_direction == "greater"
|
||||||
|
else similarity > threshold
|
||||||
|
):
|
||||||
|
return None
|
||||||
|
return similarity
|
||||||
|
|
||||||
|
# map preserves the input order, so the merge below is exactly the
|
||||||
|
# sequential reference's merge, just over precomputed scores.
|
||||||
|
with ThreadPoolExecutor(max_workers=workers) as pool:
|
||||||
|
scored = pool.map(score, records)
|
||||||
|
|
||||||
|
heap: list[_HeapEntry] = []
|
||||||
|
for record, similarity in zip(records, scored):
|
||||||
|
if similarity is None:
|
||||||
|
continue
|
||||||
|
match = SimilarityMatch(similarity, record.molecule_id, record.smiles)
|
||||||
|
rank_key = match.sort_key(threshold_direction)
|
||||||
|
entry = _HeapEntry(match, rank_key)
|
||||||
|
if len(heap) < top_k:
|
||||||
|
heapq.heappush(heap, entry)
|
||||||
|
elif rank_key < heap[0].rank_key:
|
||||||
|
heapq.heapreplace(heap, entry)
|
||||||
|
|
||||||
|
matches = [entry.match for entry in sorted(heap, key=lambda e: e.rank_key)]
|
||||||
|
return SearchResult(matches=matches, stats=stats)
|
||||||
|
|
||||||
|
|
||||||
|
def run_search_shard_parallel(
|
||||||
|
input_path: Path,
|
||||||
|
parameters: Mapping[str, object],
|
||||||
|
output_path: Path,
|
||||||
|
) -> dict[str, int]:
|
||||||
|
"""Run one planned shard with the parallel scoring core.
|
||||||
|
|
||||||
|
Accepts the same parameters as ``similarity-search`` plus ``threads``.
|
||||||
|
"""
|
||||||
|
allowed = {
|
||||||
|
"query_id",
|
||||||
|
"query_smiles",
|
||||||
|
"top_k",
|
||||||
|
"threshold",
|
||||||
|
"threshold_direction",
|
||||||
|
"progress_every",
|
||||||
|
"threads",
|
||||||
|
}
|
||||||
|
unknown = set(parameters) - allowed
|
||||||
|
if unknown:
|
||||||
|
raise ValueError(
|
||||||
|
f"unsupported similarity-search-parallel parameters: {', '.join(sorted(unknown))}"
|
||||||
|
)
|
||||||
|
query_smiles = parameters.get("query_smiles")
|
||||||
|
query_id = parameters.get("query_id")
|
||||||
|
if isinstance(query_id, str) and not isinstance(query_smiles, str):
|
||||||
|
from rdkit import Chem
|
||||||
|
|
||||||
|
from scimesh.chemistry.dataset import find_molecule_by_id
|
||||||
|
|
||||||
|
record = find_molecule_by_id(input_path, query_id)
|
||||||
|
query_smiles = Chem.MolToSmiles(record.molecule, canonical=True)
|
||||||
|
if not isinstance(query_smiles, str) or not query_smiles.strip():
|
||||||
|
raise ValueError("query_smiles is required for a distributed shard")
|
||||||
|
molecule = parse_smiles(query_smiles)
|
||||||
|
if molecule is None:
|
||||||
|
raise ValueError("query_smiles is invalid")
|
||||||
|
top_k = _positive_int(parameters.get("top_k", 20), "top_k")
|
||||||
|
threads = _nonnegative_int(parameters.get("threads", 0), "threads")
|
||||||
|
threshold = None
|
||||||
|
if "threshold" in parameters:
|
||||||
|
threshold = _unit_interval(parameters["threshold"], "threshold")
|
||||||
|
direction = parameters.get("threshold_direction", "greater")
|
||||||
|
if direction not in {"greater", "less"}:
|
||||||
|
raise ValueError("threshold_direction must be 'greater' or 'less'")
|
||||||
|
assert isinstance(direction, str)
|
||||||
|
|
||||||
|
result = search_similar_parallel(
|
||||||
|
input_path,
|
||||||
|
MoleculeRecord("query", query_smiles, molecule),
|
||||||
|
top_k=top_k,
|
||||||
|
threads=threads,
|
||||||
|
threshold=threshold,
|
||||||
|
threshold_direction=direction,
|
||||||
|
)
|
||||||
|
write_search_partial(output_path, result.matches)
|
||||||
|
return {
|
||||||
|
"scanned_rows": result.stats.scanned,
|
||||||
|
"valid_molecules": result.stats.valid,
|
||||||
|
"invalid_smiles": result.stats.invalid,
|
||||||
|
"matches_emitted": len(result.matches),
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def _positive_int(value: object, name: str) -> int:
|
||||||
|
if isinstance(value, bool) or not isinstance(value, int) or value < 1:
|
||||||
|
raise ValueError(f"{name} must be a positive integer")
|
||||||
|
return value
|
||||||
|
|
||||||
|
|
||||||
|
def _nonnegative_int(value: object, name: str) -> int:
|
||||||
|
if isinstance(value, bool) or not isinstance(value, int) or value < 0:
|
||||||
|
raise ValueError(f"{name} must be a non-negative integer")
|
||||||
|
return value
|
||||||
|
|
||||||
|
|
||||||
|
def _unit_interval(value: object, name: str) -> float:
|
||||||
|
if isinstance(value, bool) or not isinstance(value, (int, float)):
|
||||||
|
raise ValueError(f"{name} must be a number between 0 and 1")
|
||||||
|
return float(value)
|
||||||
|
|
||||||
|
|
||||||
|
# The shared shard writer is re-exported so the workload definition can reuse
|
||||||
|
# the deterministic partitioning without importing search internals.
|
||||||
|
__all__ = [
|
||||||
|
"search_similar_parallel",
|
||||||
|
"run_search_shard_parallel",
|
||||||
|
"write_search_shards",
|
||||||
|
"run_search_shard",
|
||||||
|
]
|
||||||
@@ -0,0 +1,144 @@
|
|||||||
|
"""SDK-built ``similarity-search-parallel`` workload definition and handlers.
|
||||||
|
|
||||||
|
A subclass of ``SimilaritySearchSDKWorkload``: identical contract (plan-time
|
||||||
|
query resolution, deterministic sharding, top-k reduction, byte-identical
|
||||||
|
partials), but each shard's fingerprinting and scoring runs across a thread
|
||||||
|
pool (``threads`` parameter, default = CPU count).
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from collections.abc import Mapping
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from scimesh.sdk.batch import MapReduceWorkload
|
||||||
|
from scimesh.sdk.identity import WorkloadId
|
||||||
|
from scimesh.sdk.plans import JobRequest
|
||||||
|
from scimesh.sdk.registry import WorkloadDefinition
|
||||||
|
from scimesh.sdk.ui import UIElement
|
||||||
|
|
||||||
|
from ..environment import current_environment_digest, current_scimesh_package_digest
|
||||||
|
from ..search.definition import SimilaritySearchSDKWorkload, _parameters_schema
|
||||||
|
from .core import run_search_shard_parallel
|
||||||
|
|
||||||
|
MAP_ENTRY_POINT = "scimesh.workloads.search_parallel.definition:map_search_parallel@v1"
|
||||||
|
|
||||||
|
# The parallel variant adds only the thread-count parameter on top of the
|
||||||
|
# search contract; everything else (schemas, entry points of the reduce stage,
|
||||||
|
# partitioning) is inherited.
|
||||||
|
_MAP_PARAMETERS = (
|
||||||
|
"query_id",
|
||||||
|
"query_smiles",
|
||||||
|
"top_k",
|
||||||
|
"threshold",
|
||||||
|
"threshold_direction",
|
||||||
|
"progress_every",
|
||||||
|
"threads",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _parallel_parameters_schema() -> dict[str, Any]:
|
||||||
|
schema = dict(_parameters_schema())
|
||||||
|
properties = dict(schema["properties"])
|
||||||
|
properties["threads"] = {
|
||||||
|
"type": "integer",
|
||||||
|
"minimum": 1,
|
||||||
|
"description": "Threads used to fingerprint and score one shard (default: CPU count).",
|
||||||
|
}
|
||||||
|
schema["properties"] = properties
|
||||||
|
return schema
|
||||||
|
|
||||||
|
|
||||||
|
class SimilaritySearchParallelSDKWorkload(SimilaritySearchSDKWorkload):
|
||||||
|
"""Exact top-k Tanimoto search with a per-shard thread pool."""
|
||||||
|
|
||||||
|
workload_id = WorkloadId("similarity-search-parallel", "1.0.0")
|
||||||
|
description = (
|
||||||
|
"Exact top-k Tanimoto molecular similarity search over deterministic "
|
||||||
|
"TSV shards with a bounded merge; each shard is fingerprinted and "
|
||||||
|
"scored across a thread pool. Output is byte-identical to "
|
||||||
|
"similarity-search."
|
||||||
|
)
|
||||||
|
parameters_schema = _parallel_parameters_schema()
|
||||||
|
map_parameter_names = _MAP_PARAMETERS
|
||||||
|
map_entry_point = MAP_ENTRY_POINT
|
||||||
|
ui_elements = SimilaritySearchSDKWorkload.ui_elements + (
|
||||||
|
UIElement(
|
||||||
|
"threads",
|
||||||
|
"number",
|
||||||
|
"Threads per shard",
|
||||||
|
help="Threads used to fingerprint and score one shard (default: CPU count).",
|
||||||
|
placeholder="auto",
|
||||||
|
order=6,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
def domain_validate(self, parameters: Mapping[str, Any]) -> None:
|
||||||
|
# The search base rejects unknown parameters; threads is our addition,
|
||||||
|
# so it is validated here and stripped before delegating.
|
||||||
|
rest = dict(parameters)
|
||||||
|
threads = rest.pop("threads", None)
|
||||||
|
if threads is not None and (
|
||||||
|
isinstance(threads, bool) or not isinstance(threads, int) or threads < 1
|
||||||
|
):
|
||||||
|
raise ValueError("threads must be a positive integer")
|
||||||
|
super().domain_validate(rest)
|
||||||
|
|
||||||
|
def resolved_parameters_for_plan(
|
||||||
|
self,
|
||||||
|
job,
|
||||||
|
input_path,
|
||||||
|
resolved,
|
||||||
|
):
|
||||||
|
# threads is a map-stage-only knob; strip it from the plan-level
|
||||||
|
# resolved parameters so the reduce stage projection stays clean.
|
||||||
|
resolved = super().resolved_parameters_for_plan(job, input_path, resolved)
|
||||||
|
stripped = dict(resolved)
|
||||||
|
stripped.pop("threads", None)
|
||||||
|
return stripped
|
||||||
|
|
||||||
|
def resolved_parameters(self, request: JobRequest) -> dict[str, Any]:
|
||||||
|
resolved = super().resolved_parameters(request)
|
||||||
|
if "threads" in request.parameters:
|
||||||
|
threads = request.parameters["threads"]
|
||||||
|
if isinstance(threads, bool) or not isinstance(threads, int) or threads < 1:
|
||||||
|
raise ValueError("threads must be a positive integer")
|
||||||
|
resolved["threads"] = threads
|
||||||
|
return resolved
|
||||||
|
|
||||||
|
def compute_shard(
|
||||||
|
self,
|
||||||
|
inputs: Mapping[str, Path],
|
||||||
|
parameters: Mapping[str, Any],
|
||||||
|
output_path: Path,
|
||||||
|
) -> Mapping[str, int | float]:
|
||||||
|
return run_search_shard_parallel(inputs["input"], parameters, output_path)
|
||||||
|
|
||||||
|
|
||||||
|
def similarity_search_parallel_sdk_definition(
|
||||||
|
*,
|
||||||
|
shard_rows: int = 10_000,
|
||||||
|
package_digest: str | None = None,
|
||||||
|
environment_digest: str | None = None,
|
||||||
|
) -> SimilaritySearchParallelSDKWorkload:
|
||||||
|
"""Build the SDK-built parallel similarity-search definition for tests."""
|
||||||
|
return SimilaritySearchParallelSDKWorkload(
|
||||||
|
shard_rows=shard_rows,
|
||||||
|
package_digest=package_digest or current_scimesh_package_digest(),
|
||||||
|
environment_digest=environment_digest or current_environment_digest(),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def workload_definition() -> WorkloadDefinition:
|
||||||
|
"""Installed entry-point factory for the SDK-built parallel search."""
|
||||||
|
return similarity_search_parallel_sdk_definition().definition()
|
||||||
|
|
||||||
|
|
||||||
|
def map_search_parallel(
|
||||||
|
input_path: Path,
|
||||||
|
parameters: Mapping[str, object],
|
||||||
|
output_path: Path,
|
||||||
|
) -> dict[str, int]:
|
||||||
|
"""Digest-pinned map entry point for the parallel search shard."""
|
||||||
|
return run_search_shard_parallel(input_path, parameters, output_path)
|
||||||
@@ -244,7 +244,13 @@ def test_workload_cli_exports_the_library_as_json(tmp_path: Path) -> None:
|
|||||||
assert payload["schema_version"] == 2
|
assert payload["schema_version"] == 2
|
||||||
names = [item["name"] for item in payload["workloads"]]
|
names = [item["name"] for item in payload["workloads"]]
|
||||||
assert names == sorted(
|
assert names == sorted(
|
||||||
["descriptor-batch", "molwt-filter", "similarity-graph", "similarity-search"]
|
[
|
||||||
|
"descriptor-batch",
|
||||||
|
"molwt-filter",
|
||||||
|
"similarity-graph",
|
||||||
|
"similarity-search",
|
||||||
|
"similarity-search-parallel",
|
||||||
|
]
|
||||||
)
|
)
|
||||||
for item in payload["workloads"]:
|
for item in payload["workloads"]:
|
||||||
assert item["version"] == "1.0.0"
|
assert item["version"] == "1.0.0"
|
||||||
|
|||||||
@@ -60,7 +60,7 @@ def _registered_similarity_search(shard_rows: int = 2):
|
|||||||
registry = default_sdk_registry(shard_rows=shard_rows)
|
registry = default_sdk_registry(shard_rows=shard_rows)
|
||||||
runtime = default_sdk_runtime()
|
runtime = default_sdk_runtime()
|
||||||
descriptions = registry.descriptions()
|
descriptions = registry.descriptions()
|
||||||
assert len(descriptions) == 4
|
assert len(descriptions) == 5
|
||||||
description = next(
|
description = next(
|
||||||
item for item in descriptions if item.workload.name == "similarity-search"
|
item for item in descriptions if item.workload.name == "similarity-search"
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -0,0 +1,185 @@
|
|||||||
|
"""Tests for the SDK-built similarity-search-parallel workload."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import csv
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from scimesh.sdk import (
|
||||||
|
ArtifactCollection,
|
||||||
|
DeterminismProfile,
|
||||||
|
JobRequest,
|
||||||
|
LocalArtifactStore,
|
||||||
|
LocalCoreBatchExecutor,
|
||||||
|
LocalPlanningContext,
|
||||||
|
StageKind,
|
||||||
|
)
|
||||||
|
from scimesh.workloads.library import default_sdk_registry, default_sdk_runtime
|
||||||
|
from scimesh.workloads.search.core import run_search_shard, write_search_shards
|
||||||
|
from scimesh.workloads.search_parallel import (
|
||||||
|
run_search_shard_parallel,
|
||||||
|
search_similar_parallel,
|
||||||
|
)
|
||||||
|
from scimesh.workloads.similarity_search import (
|
||||||
|
find_molecule_by_id,
|
||||||
|
search_similar,
|
||||||
|
write_search_results,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _write_dataset(path: Path, molecules: list[tuple[str, str]]) -> None:
|
||||||
|
path.write_text(
|
||||||
|
"chembl_id\tcanonical_smiles\n"
|
||||||
|
+ "".join(f"{mid}\t{smiles}\n" for mid, smiles in molecules),
|
||||||
|
encoding="utf-8",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _tie_dataset(path: Path) -> None:
|
||||||
|
# Deliberate similarity ties: propanol isomers and duplicated rows, so the
|
||||||
|
# parallel merge must reproduce the sequential row-order preference.
|
||||||
|
_write_dataset(
|
||||||
|
path,
|
||||||
|
[
|
||||||
|
("QUERY", "CCO"),
|
||||||
|
("A1", "CCCO"),
|
||||||
|
("A2", "C(CC)O"),
|
||||||
|
("B", "CCN"),
|
||||||
|
("C1", "CCC"),
|
||||||
|
("C2", "CCC"),
|
||||||
|
("D", "CCCC"),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_parallel_matches_sequential_byte_exactly(tmp_path: Path) -> None:
|
||||||
|
dataset = tmp_path / "molecules.tsv"
|
||||||
|
_tie_dataset(dataset)
|
||||||
|
query = find_molecule_by_id(dataset, "QUERY")
|
||||||
|
|
||||||
|
reference = search_similar(dataset, query, top_k=5, progress_every=0)
|
||||||
|
reference_path = tmp_path / "reference.csv"
|
||||||
|
write_search_results(reference_path, reference.matches)
|
||||||
|
|
||||||
|
for threads in (1, 2, 4):
|
||||||
|
parallel = search_similar_parallel(dataset, query, top_k=5, threads=threads)
|
||||||
|
parallel_path = tmp_path / f"parallel-{threads}.csv"
|
||||||
|
write_search_results(parallel_path, parallel.matches)
|
||||||
|
assert parallel_path.read_bytes() == reference_path.read_bytes(), (
|
||||||
|
f"threads={threads} diverged from the reference"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_parallel_shard_matches_sequential_shard(tmp_path: Path) -> None:
|
||||||
|
dataset = tmp_path / "molecules.tsv"
|
||||||
|
_tie_dataset(dataset)
|
||||||
|
shard_dir = tmp_path / "shards"
|
||||||
|
shard_dir.mkdir()
|
||||||
|
shards = write_search_shards(dataset, shard_dir, shard_rows=2)
|
||||||
|
parameters = {"query_smiles": "CCO", "top_k": 3, "threads": 4}
|
||||||
|
|
||||||
|
sequential_out = tmp_path / "seq.tsv"
|
||||||
|
parallel_out = tmp_path / "par.tsv"
|
||||||
|
sequential_parameters = dict(parameters)
|
||||||
|
sequential_parameters.pop("threads")
|
||||||
|
run_search_shard(shards[0], sequential_parameters, sequential_out)
|
||||||
|
run_search_shard_parallel(shards[0], parameters, parallel_out)
|
||||||
|
assert parallel_out.read_bytes() == sequential_out.read_bytes()
|
||||||
|
|
||||||
|
with parallel_out.open(encoding="utf-8") as handle:
|
||||||
|
rows = list(csv.DictReader(handle))
|
||||||
|
assert rows[0]["rank"] == "1"
|
||||||
|
assert rows[0]["similarity"].startswith("0.5") # CCO vs CCCO
|
||||||
|
assert len(rows) <= 3
|
||||||
|
|
||||||
|
|
||||||
|
def test_parallel_rejects_bad_parameters(tmp_path: Path) -> None:
|
||||||
|
dataset = tmp_path / "molecules.tsv"
|
||||||
|
_tie_dataset(dataset)
|
||||||
|
output = tmp_path / "out.tsv"
|
||||||
|
|
||||||
|
with pytest.raises(ValueError, match="threads must be a non-negative integer"):
|
||||||
|
run_search_shard_parallel(
|
||||||
|
dataset, {"query_smiles": "CCO", "threads": -1}, output
|
||||||
|
)
|
||||||
|
with pytest.raises(ValueError, match="unsupported"):
|
||||||
|
run_search_shard_parallel(dataset, {"query_smiles": "CCO", "nope": 1}, output)
|
||||||
|
with pytest.raises(ValueError, match="query_smiles is invalid"):
|
||||||
|
run_search_shard_parallel(dataset, {"query_smiles": "СС"}, output)
|
||||||
|
|
||||||
|
|
||||||
|
def _registered_parallel_search(shard_rows: int = 2):
|
||||||
|
registry = default_sdk_registry(shard_rows=shard_rows)
|
||||||
|
runtime = default_sdk_runtime()
|
||||||
|
description = next(
|
||||||
|
item
|
||||||
|
for item in registry.descriptions()
|
||||||
|
if item.workload.name == "similarity-search-parallel"
|
||||||
|
)
|
||||||
|
definition, negotiated = registry.require(
|
||||||
|
description.workload.name,
|
||||||
|
description.workload.version,
|
||||||
|
description.package_digest,
|
||||||
|
runtime=runtime,
|
||||||
|
)
|
||||||
|
return registry, runtime, description, definition, negotiated
|
||||||
|
|
||||||
|
|
||||||
|
def test_parallel_manifest_is_registered_and_negotiable() -> None:
|
||||||
|
_, runtime, description, definition, negotiated = _registered_parallel_search()
|
||||||
|
manifest = definition.manifest
|
||||||
|
|
||||||
|
assert description.enabled is True
|
||||||
|
assert manifest.workload.name == "similarity-search-parallel"
|
||||||
|
assert manifest.workload.version == "1.0.0"
|
||||||
|
assert manifest.determinism is DeterminismProfile.BYTE_EXACT
|
||||||
|
assert manifest.verifier.verifier.canonical == "exact-artifact@1"
|
||||||
|
assert set(mode.value for mode in manifest.trust_modes) == {
|
||||||
|
"trusted",
|
||||||
|
"untrusted_quorum",
|
||||||
|
}
|
||||||
|
assert [stage.kind for stage in manifest.workflow.stages] == [
|
||||||
|
StageKind.MAP,
|
||||||
|
StageKind.REDUCE,
|
||||||
|
]
|
||||||
|
assert "threads" in manifest.parameters_schema["properties"]
|
||||||
|
assert negotiated is not None
|
||||||
|
assert runtime is not None
|
||||||
|
|
||||||
|
|
||||||
|
def test_parallel_executor_matches_reference(tmp_path: Path) -> None:
|
||||||
|
dataset = tmp_path / "molecules.tsv"
|
||||||
|
_tie_dataset(dataset)
|
||||||
|
registry, runtime, description, definition, _ = _registered_parallel_search()
|
||||||
|
artifact_store = LocalArtifactStore(tmp_path / "artifacts")
|
||||||
|
input_port = definition.manifest.inputs["input"]
|
||||||
|
dataset_artifact = artifact_store.import_file(
|
||||||
|
dataset,
|
||||||
|
declaration=input_port.schema,
|
||||||
|
)
|
||||||
|
request = JobRequest(
|
||||||
|
workload=definition.manifest.workload,
|
||||||
|
parameters={"query_id": "QUERY", "top_k": 3, "threads": 2},
|
||||||
|
inputs={"input": ArtifactCollection.single(dataset_artifact)},
|
||||||
|
)
|
||||||
|
|
||||||
|
result = LocalCoreBatchExecutor(
|
||||||
|
registry,
|
||||||
|
runtime,
|
||||||
|
artifact_store,
|
||||||
|
tmp_path / "sdk-work",
|
||||||
|
).execute(request, description.package_digest)
|
||||||
|
result_artifact = result.outputs["result"].items[0].artifact
|
||||||
|
|
||||||
|
reference_path = tmp_path / "reference.csv"
|
||||||
|
query = find_molecule_by_id(dataset, "QUERY")
|
||||||
|
reference = search_similar(dataset, query, top_k=3, progress_every=0)
|
||||||
|
write_search_results(reference_path, reference.matches)
|
||||||
|
|
||||||
|
assert (
|
||||||
|
artifact_store.materialize(result_artifact).read_bytes()
|
||||||
|
== reference_path.read_bytes()
|
||||||
|
)
|
||||||
|
assert result.task_key == "reduce/final"
|
||||||
Reference in New Issue
Block a user