Compare commits
11
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
ecc8944006 | ||
|
|
dc15e2d04b | ||
|
|
ff1fc25d77 | ||
|
|
6b67326c3b | ||
|
|
d57f8778ac | ||
|
|
5b2d5b5f6e | ||
|
|
b63441bb8c | ||
|
|
91d049b4df | ||
|
|
70c9e567ca | ||
|
|
81e854eb1d | ||
|
|
9a90d59ecf |
@@ -38,8 +38,36 @@
|
||||
## Plan (предыдущая задача — выполнена)
|
||||
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 (ночная сессия)
|
||||
- (начало) Пустое имя воркера: валидация добавлена в `domain.NewWorker`, добавлен `TestNewWorkerRejectsBlankName`; в `worker_test.go` сломан тест из-за `fixedTime` vs `testNow` — на паузе, продолжить.
|
||||
- ✅ **П.1 Пустое имя воркера**: `domain.NewWorker` нормализует/отклоняет пустое имя + `TestNewWorkerRejectsBlankName` (починен `fixedTime`→`testNow`).
|
||||
- ✅ **П.2 `--check`**: пробует managed venv (если установлен) + реальный пробинг учётки (exchange ключа / claim-пробa) — `CheckAuth` + тесты; на машине пользователя: `✓ auth: credential accepted`, venv python, scimesh installed.
|
||||
- ✅ **П.3 Рестайлинг**: единый CSS-partial `ui-base.html` (дизайн-система админки), все 5 страниц (new-job, job, workloads, add-worker, profile) переведены, проверены в браузере без console-ошибок.
|
||||
- ✅ **П.4 Postgres integration**: admin-методы (SetTrust, ListJobsPaginated, TaskCounts, byDay/byWorkload, TaskStats, ArtifactSize, DB size, WorkloadSettings) + `ensureMigrated` для порядка запуска; весь suite зелёный.
|
||||
- ✅ **П.5 Статистика воркера в визарде**: лог `task claimed` в агента + парсинг registered/claimed/completed/failed → `/api/status.stats` + карточки в статусной странице + тест.
|
||||
- ✅ **П.6 Prune артефактов**: `JobRepository.ListCompletedBefore/Delete` (sqlite+postgres+memstore), usecase `PruneArtifacts` (каскад + blob-файлы), `POST /ui/admin/api/prune`, кнопка в Settings, тесты (sqlite+usecase); E2E: 200, freed bytes.
|
||||
- ✅ **П.7 Удаление offline-воркеров**: `WorkerRepository.Delete` + `Admin.RemoveWorker` (только offline) + `POST /ui/admin/api/workers/{id}/remove` + кнопка в Workers + тест; E2E: 204, строка удалена.
|
||||
- ✅ **П.8 setuptools_scm**: `dynamic = ["version"]`, CI-джоба wheel без sed (fetch-depth 0); локальная проверка: wheel на теге = `scimesh-1.1.0a16-py3-none-any.whl` (совпадает с Go-нормализацией); релизный ассет подтверждён.
|
||||
- ✅ **П.9 Docs**: STATUS.md синхронизирован (админка, визард, wheel, CTX-19/20 implemented).
|
||||
- ✅ **П.10 (доп.) Баг в serve-режиме**: worker-key exchange был недоступен снаружи (userservice на loopback) — добавлен прокси `POST /worker-tokens/exchange` на координаторе, `PublicUserserviceURL=""` + fallback на origin в add-worker. Проверено E2E.
|
||||
- ✅ **П.10 Quorum E2E (Docker)**: 2 untrusted-воркера с разными ключами (alice/bob) → джоб completed 3/3, в task_results по 2 голоса от разных владельцев с одинаковым sha256 → результат байт-в-байт = локальному эталону.
|
||||
- ✅ **Финальный гейт**: `go test -race ./...` ✅, golangci-lint 0 issues ✅, pytest 208 ✅, postgres integration ✅, Windows кросс-сборка ✅.
|
||||
- ✅ **Релиз v1.1.0-alpha.16** (бинарники + wheel `scimesh-1.1.0a16`), все воркфлоу success; бинарники на машине пользователя обновлены до alpha.16.
|
||||
|
||||
## Progress (прошлая работа — выполнено)
|
||||
- ✅ Релизы alpha.12–15: фикс версии визарда, venv task_runner, preflight через venv, MkdirAll при скачивании wheel, кнопка Install в шаблоне.
|
||||
@@ -47,5 +75,8 @@
|
||||
- ✅ На машине пользователя: визард alpha.15, правильный токен, venv из wheel, воркер emil-pc online, 15 пустых воркеров вычищены из БД.
|
||||
- ✅ Гейт: race + lint + pytest 208.
|
||||
|
||||
## Completion
|
||||
COMPLETED — ночной план выполнен полностью (10 пунктов + 2 найденных бага, включая E2E quorum на релизном коде). Все гейты зелёные, релиз v1.1.0-alpha.16 опубликован.
|
||||
|
||||
## Completion (предыдущая задача)
|
||||
COMPLETED — пайплайн доведён до рабочего состояния и проверен на релизных артефактах v1.1.0-alpha.14.
|
||||
|
||||
@@ -10,11 +10,13 @@ import (
|
||||
"errors"
|
||||
"flag"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"os"
|
||||
"os/exec"
|
||||
"os/signal"
|
||||
"strings"
|
||||
"syscall"
|
||||
"time"
|
||||
|
||||
@@ -166,6 +168,16 @@ func runSetup(args []string) int {
|
||||
})
|
||||
listener, err := server.Listen()
|
||||
if err != nil {
|
||||
// A wizard may already be running on this port (left open, or a
|
||||
// second terminal). If it is ours, opening the browser is the
|
||||
// friendlier outcome than failing the command.
|
||||
if existing := wizardAlreadyRunning(*port); existing != "" {
|
||||
logger.Info("the setup wizard is already running", "url", existing)
|
||||
if !*noOpen {
|
||||
openBrowser(existing)
|
||||
}
|
||||
return 0
|
||||
}
|
||||
logger.Error("setup wizard could not bind the loopback port", "err", err)
|
||||
return 1
|
||||
}
|
||||
@@ -199,3 +211,28 @@ func openBrowser(url string) {
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
// wizardAlreadyRunning probes the requested loopback port and returns its URL
|
||||
// when it serves the setup wizard page, or "" when it does not (another
|
||||
// process, or nothing at all).
|
||||
func wizardAlreadyRunning(port int) string {
|
||||
url := fmt.Sprintf("http://127.0.0.1:%d/", port)
|
||||
client := http.Client{Timeout: 2 * time.Second}
|
||||
req, err := http.NewRequestWithContext(context.Background(), http.MethodGet, url, nil)
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
defer func() { _ = resp.Body.Close() }()
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return ""
|
||||
}
|
||||
body, err := io.ReadAll(io.LimitReader(resp.Body, 1<<16))
|
||||
if err != nil || !strings.Contains(string(body), "SciMesh Worker · Setup") {
|
||||
return ""
|
||||
}
|
||||
return url
|
||||
}
|
||||
|
||||
@@ -88,8 +88,10 @@ func CheckEnvironment(ctx context.Context) CheckReport {
|
||||
// execute with.
|
||||
func CheckEnvironmentWithPython(ctx context.Context, python string) CheckReport {
|
||||
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.
|
||||
//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, "-c", "import importlib.metadata as m; print(m.version('scimesh'))")
|
||||
out, err := cmd.Output()
|
||||
if err != nil {
|
||||
// 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
|
||||
MaxTasks int // 0 = unlimited
|
||||
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) {
|
||||
@@ -145,9 +149,22 @@ func LoadConfig() (*Config, error) {
|
||||
TaskRunner: runner,
|
||||
MaxTasks: maxTasks,
|
||||
ExitWhenIdle: os.Getenv("EXIT_WHEN_IDLE") == "1",
|
||||
Concurrency: envInt("WORKER_CONCURRENCY", 1),
|
||||
}, 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) {
|
||||
raw := os.Getenv(name)
|
||||
if raw == "" {
|
||||
|
||||
@@ -21,6 +21,7 @@ type ConfigFile struct {
|
||||
WorkerName string `json:"worker_name,omitempty"`
|
||||
CPUCount int `json:"cpu_count"`
|
||||
MemoryMB int `json:"memory_mb"`
|
||||
Concurrency int `json:"concurrency,omitempty"`
|
||||
TaskRunner []string `json:"task_runner,omitempty"`
|
||||
}
|
||||
|
||||
@@ -103,6 +104,10 @@ func (f *ConfigFile) Config() (*Config, error) {
|
||||
if config.MemoryMB < 0 {
|
||||
config.MemoryMB = 0
|
||||
}
|
||||
config.Concurrency = f.Concurrency
|
||||
if config.Concurrency < 1 {
|
||||
config.Concurrency = 1
|
||||
}
|
||||
if len(f.TaskRunner) > 0 {
|
||||
config.TaskRunner = f.TaskRunner
|
||||
}
|
||||
|
||||
@@ -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}
|
||||
}
|
||||
|
||||
// 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 {
|
||||
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
|
||||
for {
|
||||
if !d.registered {
|
||||
if err := d.register(); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
d.cleanupExpiredDirectories()
|
||||
outcome, err := d.runOnce()
|
||||
if err != nil {
|
||||
@@ -63,8 +100,11 @@ func (d *Daemon) RunForever() error {
|
||||
}
|
||||
failures = 0
|
||||
if outcome.Claimed && outcome.Completed {
|
||||
d.mu.Lock()
|
||||
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")
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -11,6 +11,7 @@ import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
@@ -55,6 +56,7 @@ type fakeCoordinator struct {
|
||||
uploadSize int64
|
||||
inputBytes []byte
|
||||
conflict bool // 409 on heartbeat/upload/result
|
||||
mu sync.Mutex
|
||||
}
|
||||
|
||||
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))
|
||||
var server *httptest.Server
|
||||
server = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
fake.mu.Lock()
|
||||
defer fake.mu.Unlock()
|
||||
switch {
|
||||
case r.Method == http.MethodPost && r.URL.Path == "/workers/register":
|
||||
writeJSON(w, http.StatusCreated, map[string]any{
|
||||
@@ -294,3 +298,78 @@ func TestDaemonIdleClaimIsNotCompleted(t *testing.T) {
|
||||
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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -377,6 +377,7 @@ type saveConfigRequest struct {
|
||||
WorkerName string `json:"worker_name"`
|
||||
CPUCount int `json:"cpu_count"`
|
||||
MemoryMB int `json:"memory_mb"`
|
||||
Concurrency int `json:"concurrency"`
|
||||
TaskRunner []string `json:"task_runner"`
|
||||
}
|
||||
|
||||
@@ -395,6 +396,7 @@ func (s *Server) handleSaveConfig(w http.ResponseWriter, r *http.Request) {
|
||||
WorkerName: strings.TrimSpace(req.WorkerName),
|
||||
CPUCount: req.CPUCount,
|
||||
MemoryMB: req.MemoryMB,
|
||||
Concurrency: req.Concurrency,
|
||||
TaskRunner: req.TaskRunner,
|
||||
}
|
||||
if file.CoordinatorURL == "" {
|
||||
@@ -418,6 +420,9 @@ func (s *Server) handleSaveConfig(w http.ResponseWriter, r *http.Request) {
|
||||
if 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;
|
||||
// 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.
|
||||
@@ -449,6 +454,9 @@ func (s *Server) handleTest(w http.ResponseWriter, r *http.Request) {
|
||||
// checking the bare system python3 would keep reporting scimesh as
|
||||
// missing even though the worker would run with the venv.
|
||||
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)
|
||||
}
|
||||
|
||||
@@ -623,3 +631,26 @@ func truncate(s string, n int) string {
|
||||
}
|
||||
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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
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 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>
|
||||
</div>
|
||||
@@ -228,7 +229,8 @@ function draftConfig(){
|
||||
userservice_url:state.mode==='key'?$('in-users').value.trim():'',
|
||||
work_dir:$('in-dir').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'];
|
||||
return cfg;
|
||||
|
||||
@@ -597,6 +597,206 @@
|
||||
"upload_ready": true,
|
||||
"verifier": "exact-artifact@1",
|
||||
"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"]
|
||||
"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"
|
||||
"descriptor-batch@1.0.0" = "scimesh.workloads.descriptors: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 .molwt_filter import molwt_filter_sdk_definition
|
||||
from .search import similarity_search_sdk_definition
|
||||
from .search_parallel import similarity_search_parallel_sdk_definition
|
||||
|
||||
__all__ = [
|
||||
"default_sdk_registry",
|
||||
@@ -46,6 +47,10 @@ def default_sdk_registry(
|
||||
similarity_search_sdk_definition(shard_rows=shard_rows).definition(),
|
||||
enabled=True,
|
||||
)
|
||||
registry.register(
|
||||
similarity_search_parallel_sdk_definition(shard_rows=shard_rows).definition(),
|
||||
enabled=True,
|
||||
)
|
||||
registry.register(
|
||||
similarity_graph_sdk_definition().definition(),
|
||||
enabled=True,
|
||||
@@ -82,6 +87,7 @@ def default_sdk_runtime(
|
||||
workload_capabilities
|
||||
or (
|
||||
"similarity-search",
|
||||
"similarity-search-parallel",
|
||||
"similarity-graph",
|
||||
"descriptor-batch",
|
||||
"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
|
||||
names = [item["name"] for item in payload["workloads"]]
|
||||
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"]:
|
||||
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)
|
||||
runtime = default_sdk_runtime()
|
||||
descriptions = registry.descriptions()
|
||||
assert len(descriptions) == 4
|
||||
assert len(descriptions) == 5
|
||||
description = next(
|
||||
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