mirror of
https://github.com/gadgethd/ukmesh.git
synced 2026-09-02 20:44:17 +00:00
1746 lines
67 KiB
Python
1746 lines
67 KiB
Python
"""
|
|
ML path learner worker.
|
|
|
|
Extracts gold-standard path observations from uniquely-resolved multibyte
|
|
packets, trains a LightGBM model to predict which node a 1-byte hash maps to,
|
|
and writes high-confidence scores back to ml_path_prefix_scores for the lazy
|
|
resolver to use as its highest-priority evidence tier.
|
|
|
|
Training data source:
|
|
Multibyte packets (path_hash_size_bytes > 1) where every hop hash uniquely
|
|
resolves to exactly one positioned node are treated as ground truth. Their
|
|
hashes are degraded to 2-char (1-byte) prefixes to simulate the hard case.
|
|
|
|
Feedback loop prevention:
|
|
Only multibyte packet resolutions (never lazy-resolver output) are used as
|
|
labels. Model predictions are never fed back as training data.
|
|
"""
|
|
|
|
import io
|
|
import json
|
|
import logging
|
|
import math
|
|
import os
|
|
import random
|
|
import signal
|
|
import threading
|
|
import time
|
|
import uuid
|
|
import warnings
|
|
from collections import defaultdict
|
|
from dataclasses import dataclass
|
|
from datetime import datetime, timezone
|
|
|
|
import joblib
|
|
import numpy as np
|
|
import psycopg2
|
|
import psycopg2.extras
|
|
from lightgbm import LGBMClassifier
|
|
from sklearn.calibration import CalibratedClassifierCV
|
|
from sklearn.frozen import FrozenEstimator
|
|
|
|
# LightGBM 4.5 calls sklearn's internal check_array with the old
|
|
# force_all_finite= kwarg; suppress until LightGBM is updated.
|
|
warnings.filterwarnings('ignore', message=".*force_all_finite.*", category=FutureWarning)
|
|
warnings.filterwarnings('ignore', message=".*force_all_finite.*", category=UserWarning)
|
|
|
|
logging.basicConfig(
|
|
level=logging.INFO,
|
|
format='%(asctime)s [ml-learner] %(levelname)s %(message)s',
|
|
)
|
|
log = logging.getLogger(__name__)
|
|
|
|
DATABASE_URL = os.environ.get('DATABASE_URL')
|
|
GOLD_INTERVAL_SECS = int(os.environ.get('GOLD_EXTRACTION_INTERVAL_MINS', '15')) * 60
|
|
TRAIN_INTERVAL_SECS = int(os.environ.get('TRAINING_INTERVAL_MINS', '30')) * 60
|
|
MIN_GOLD_ROWS = int(os.environ.get('MIN_GOLD_PATHS', '100'))
|
|
CONFIDENCE_THRESHOLD = float(os.environ.get('CONFIDENCE_THRESHOLD', '0.85'))
|
|
MIN_OBSERVATION_COUNT = int(os.environ.get('MIN_OBSERVATION_COUNT', '1'))
|
|
PROMOTION_MIN_DELTA = float(os.environ.get('PROMOTION_MIN_DELTA', '0.0'))
|
|
RETAIN_VARIANT_RESULT_GENERATIONS = int(os.environ.get('ML_RETAIN_VARIANT_RESULT_GENERATIONS', '96'))
|
|
RETAIN_MODEL_ARTIFACT_GENERATIONS = int(os.environ.get('ML_RETAIN_MODEL_ARTIFACT_GENERATIONS', '96'))
|
|
CLEANUP_GENERATION_BATCH_SIZE = int(os.environ.get('ML_CLEANUP_GENERATION_BATCH_SIZE', '24'))
|
|
MAX_HOP_KM = 150.0
|
|
GOLD_BATCH = 5000
|
|
MAX_TRAINING_GOLD_ROWS = max(MIN_GOLD_ROWS, int(os.environ.get('MAX_TRAINING_GOLD_ROWS', '100000')))
|
|
CHECKPOINT_KEY = 'gold_extraction_checkpoint'
|
|
ML_MODEL_VERSION = os.environ.get('ML_MODEL_VERSION', 'lightgbm-path-v1')
|
|
ML_DATA_VERSION = os.environ.get('ML_DATA_VERSION', 'gold-multibyte-v2')
|
|
ML_LEASE_SECONDS = max(30, min(600, int(os.environ.get('ML_LEASE_SECONDS', '90'))))
|
|
ML_RUN_DEADLINE_SECONDS = max(60, min(6 * 3600, int(os.environ.get('ML_RUN_DEADLINE_SECONDS', '1200'))))
|
|
ML_HEARTBEAT_SECONDS = max(5, min(ML_LEASE_SECONDS // 3, int(os.environ.get('ML_HEARTBEAT_SECONDS', '20'))))
|
|
STOP_EVENT = threading.Event()
|
|
|
|
# ── Genetic / evolutionary search ────────────────────────────────────────────
|
|
|
|
POPULATION_SIZE = int(os.environ.get('POPULATION_SIZE', '10'))
|
|
RANDOM_SEED = int(os.environ.get('RANDOM_SEED', '42'))
|
|
|
|
DEFAULT_PARAMS: dict = {
|
|
'num_leaves': 31,
|
|
'learning_rate': 0.05,
|
|
'min_child_samples': 20,
|
|
'n_estimators': 300,
|
|
'feature_fraction': 1.0,
|
|
'bagging_fraction': 1.0,
|
|
'bagging_freq': 0,
|
|
}
|
|
|
|
PARAM_BOUNDS: dict = {
|
|
'num_leaves': (8, 127),
|
|
'learning_rate': (0.01, 0.20),
|
|
'min_child_samples': (3, 50),
|
|
'n_estimators': (100, 600),
|
|
'feature_fraction': (0.5, 1.0),
|
|
'bagging_fraction': (0.5, 1.0),
|
|
'bagging_freq': (0, 5),
|
|
}
|
|
|
|
INT_PARAMS = {'num_leaves', 'min_child_samples', 'n_estimators', 'bagging_freq'}
|
|
|
|
|
|
def _clamp_param(key: str, value: float) -> int | float:
|
|
lo, hi = PARAM_BOUNDS[key]
|
|
if key in INT_PARAMS:
|
|
return max(lo, min(hi, int(round(value))))
|
|
return max(lo, min(hi, float(value)))
|
|
|
|
|
|
def random_params() -> dict:
|
|
"""Return a fully random hyperparameter set inside the allowed bounds."""
|
|
return {
|
|
key: random.randint(lo, hi) if key in INT_PARAMS else random.uniform(lo, hi)
|
|
for key, (lo, hi) in PARAM_BOUNDS.items()
|
|
}
|
|
|
|
|
|
def mutate(params: dict, strength: float = 0.35) -> dict:
|
|
"""Return a new param dict with random perturbations from params."""
|
|
new: dict = {}
|
|
for k, v in params.items():
|
|
if random.random() < 0.85:
|
|
factor = random.uniform(1.0 - strength, 1.0 + strength)
|
|
raw = v * factor
|
|
new[k] = _clamp_param(k, raw)
|
|
else:
|
|
new[k] = v
|
|
return new
|
|
|
|
|
|
def create_population(base_params: dict, generation: int = 1) -> list[dict]:
|
|
"""Build a mixed population seeded from the champion plus wider exploration."""
|
|
exploration = min(1.0, 0.30 + (max(0, generation - 1) * 0.08))
|
|
strengths = [0.25, 0.45, 0.70, exploration]
|
|
pop = [base_params.copy()]
|
|
seen = {tuple(sorted(base_params.items()))}
|
|
while len(pop) < POPULATION_SIZE:
|
|
if len(pop) % 4 == 0:
|
|
candidate = random_params()
|
|
else:
|
|
candidate = mutate(base_params, random.choice(strengths))
|
|
key = tuple(sorted(candidate.items()))
|
|
if key in seen:
|
|
continue
|
|
seen.add(key)
|
|
pop.append(candidate)
|
|
return pop
|
|
|
|
|
|
GLOBAL_NETWORK = 'global'
|
|
COMBINED_UKMESH_NETWORKS = ('teesside', 'ukmesh')
|
|
COMBINED_UKMESH_SCOPE = 'ukmesh_combined'
|
|
|
|
|
|
def network_scope_key(network: str) -> str:
|
|
return COMBINED_UKMESH_SCOPE if network in COMBINED_UKMESH_NETWORKS else network
|
|
|
|
|
|
def network_scope_values(network_or_scope: str) -> list[str]:
|
|
if network_or_scope == COMBINED_UKMESH_SCOPE or network_or_scope in COMBINED_UKMESH_NETWORKS:
|
|
return list(COMBINED_UKMESH_NETWORKS)
|
|
return [network_or_scope]
|
|
|
|
|
|
def get_champion_params(db) -> dict:
|
|
"""Return hyperparams of the current global champion, or DEFAULT_PARAMS."""
|
|
with db.cursor() as cur:
|
|
cur.execute(
|
|
"""SELECT hyperparams FROM ml_model_versions
|
|
WHERE network = %s AND is_active = TRUE
|
|
ORDER BY promoted_at DESC LIMIT 1""",
|
|
[GLOBAL_NETWORK],
|
|
)
|
|
row = cur.fetchone()
|
|
if row and row['hyperparams']:
|
|
p = row['hyperparams']
|
|
return {k: p.get(k, v) for k, v in DEFAULT_PARAMS.items()}
|
|
return DEFAULT_PARAMS.copy()
|
|
|
|
|
|
def get_current_generation(db) -> int:
|
|
"""Return the latest global generation number (0 if none)."""
|
|
with db.cursor() as cur:
|
|
cur.execute(
|
|
"""SELECT COALESCE(MAX(generation), 0) AS generation
|
|
FROM (
|
|
SELECT generation
|
|
FROM ml_model_versions
|
|
WHERE network = %s
|
|
UNION ALL
|
|
SELECT generation
|
|
FROM ml_model_variant_runs
|
|
WHERE model_network = %s
|
|
) generations""",
|
|
[GLOBAL_NETWORK, GLOBAL_NETWORK],
|
|
)
|
|
row = cur.fetchone()
|
|
return int(row['generation']) if row and row['generation'] is not None else 0
|
|
|
|
|
|
# ── Geometry ──────────────────────────────────────────────────────────────────
|
|
|
|
def dist_km(lat1: float, lon1: float, lat2: float, lon2: float) -> float:
|
|
R = 6371.0
|
|
d_lat = math.radians(lat2 - lat1)
|
|
d_lon = math.radians(lon2 - lon1)
|
|
a = math.sin(d_lat / 2) ** 2 + math.cos(math.radians(lat1)) * math.cos(math.radians(lat2)) * math.sin(d_lon / 2) ** 2
|
|
return 2 * R * math.asin(math.sqrt(a))
|
|
|
|
|
|
# ── Database connection ───────────────────────────────────────────────────────
|
|
|
|
def get_db():
|
|
if not DATABASE_URL:
|
|
raise RuntimeError('DATABASE_URL is required')
|
|
conn = psycopg2.connect(DATABASE_URL, cursor_factory=psycopg2.extras.RealDictCursor)
|
|
conn.autocommit = False
|
|
return conn
|
|
|
|
|
|
class LeaseLost(RuntimeError):
|
|
pass
|
|
|
|
|
|
class RunDeadlineExceeded(RuntimeError):
|
|
pass
|
|
|
|
|
|
class RunGuard:
|
|
def __init__(
|
|
self,
|
|
token: str,
|
|
deadline_monotonic: float,
|
|
should_train: bool = False,
|
|
expected_model_version: str = ML_MODEL_VERSION,
|
|
):
|
|
self.token = token
|
|
self.deadline_monotonic = deadline_monotonic
|
|
self.should_train = should_train
|
|
self.expected_model_version = expected_model_version
|
|
self.stop_event = threading.Event()
|
|
self.lease_lost = threading.Event()
|
|
self.thread = threading.Thread(
|
|
target=self._heartbeat,
|
|
name=f'ml-lease-{token[:8]}',
|
|
daemon=True,
|
|
)
|
|
|
|
def _heartbeat(self):
|
|
db = None
|
|
while not self.stop_event.wait(ML_HEARTBEAT_SECONDS):
|
|
try:
|
|
if db is None or db.closed:
|
|
db = get_db()
|
|
with db.cursor() as cur:
|
|
cur.execute(
|
|
"""UPDATE ml_learner_state
|
|
SET heartbeat_at = NOW(),
|
|
lease_expires_at = NOW() + (%s * INTERVAL '1 second'),
|
|
updated_at = NOW()
|
|
WHERE singleton = TRUE
|
|
AND leader_token = %s
|
|
AND run_deadline_at > NOW()
|
|
RETURNING singleton""",
|
|
[ML_LEASE_SECONDS, self.token],
|
|
)
|
|
owned = cur.fetchone() is not None
|
|
db.commit()
|
|
if not owned:
|
|
self.lease_lost.set()
|
|
break
|
|
except Exception as exc:
|
|
log.warning('ML lease heartbeat reconnecting: %s', exc)
|
|
try:
|
|
if db is not None and not db.closed:
|
|
db.rollback()
|
|
db.close()
|
|
except Exception:
|
|
pass
|
|
db = None
|
|
if db is not None and not db.closed:
|
|
db.close()
|
|
|
|
def check(self):
|
|
if STOP_EVENT.is_set():
|
|
raise InterruptedError('ML learner is shutting down')
|
|
if self.lease_lost.is_set():
|
|
raise LeaseLost('ML leader lease was lost')
|
|
if time.monotonic() >= self.deadline_monotonic:
|
|
raise RunDeadlineExceeded('ML run deadline exceeded')
|
|
|
|
def fence_publication(self, cursor, published_version: str):
|
|
self.check()
|
|
cursor.execute(
|
|
"""UPDATE ml_learner_state
|
|
SET model_version = %s,
|
|
data_version = %s,
|
|
heartbeat_at = NOW(),
|
|
updated_at = NOW()
|
|
WHERE singleton = TRUE
|
|
AND leader_token = %s
|
|
AND model_version = %s
|
|
AND lease_expires_at > NOW()
|
|
AND run_deadline_at > NOW()
|
|
RETURNING singleton""",
|
|
[
|
|
published_version,
|
|
ML_DATA_VERSION,
|
|
self.token,
|
|
self.expected_model_version,
|
|
],
|
|
)
|
|
if cursor.fetchone() is None:
|
|
self.lease_lost.set()
|
|
raise LeaseLost('ML publication fence rejected stale leader')
|
|
|
|
def fence_checkpoint(self, cursor, checkpoint: 'ExtractionCursor'):
|
|
self.check()
|
|
cursor.execute(
|
|
"""UPDATE ml_learner_state
|
|
SET cursor_observed_at = %s,
|
|
cursor_packet_hash = %s,
|
|
cursor_network = %s,
|
|
cursor_rx_node_id = %s,
|
|
cursor_topic = %s,
|
|
cursor_raw_hex = %s,
|
|
heartbeat_at = NOW(),
|
|
updated_at = NOW()
|
|
WHERE singleton = TRUE
|
|
AND leader_token = %s
|
|
AND lease_expires_at > NOW()
|
|
AND run_deadline_at > NOW()
|
|
RETURNING singleton""",
|
|
[*checkpoint.sql_values(), self.token],
|
|
)
|
|
if cursor.fetchone() is None:
|
|
self.lease_lost.set()
|
|
raise LeaseLost('ML checkpoint fence rejected stale leader')
|
|
|
|
def finish(
|
|
self,
|
|
db,
|
|
reason: str,
|
|
training_completed: bool = False,
|
|
retry_delay_seconds: int = GOLD_INTERVAL_SECS,
|
|
):
|
|
self.stop_event.set()
|
|
if self.thread.is_alive():
|
|
self.thread.join(timeout=5)
|
|
with db.cursor() as cur:
|
|
cur.execute(
|
|
"""UPDATE ml_learner_state
|
|
SET leader_token = NULL,
|
|
lease_expires_at = NULL,
|
|
heartbeat_at = NOW(),
|
|
run_deadline_at = NULL,
|
|
next_run_at = NOW() + (%s * INTERVAL '1 second'),
|
|
last_trained_at = CASE
|
|
WHEN %s THEN NOW()
|
|
ELSE last_trained_at
|
|
END,
|
|
next_training_at = CASE
|
|
WHEN %s THEN NOW() + (%s * INTERVAL '1 second')
|
|
ELSE next_training_at
|
|
END,
|
|
last_terminal_reason = %s,
|
|
updated_at = NOW()
|
|
WHERE singleton = TRUE AND leader_token = %s""",
|
|
[
|
|
max(0, retry_delay_seconds),
|
|
training_completed,
|
|
training_completed,
|
|
TRAIN_INTERVAL_SECS,
|
|
reason[:240],
|
|
self.token,
|
|
],
|
|
)
|
|
db.commit()
|
|
|
|
def __enter__(self):
|
|
self.thread.start()
|
|
return self
|
|
|
|
def __exit__(self, exc_type, exc, tb):
|
|
self.stop_event.set()
|
|
self.thread.join(timeout=5)
|
|
|
|
|
|
def claim_leadership(db) -> RunGuard | None:
|
|
token = uuid.uuid4().hex
|
|
with db.cursor() as cur:
|
|
cur.execute(
|
|
"""UPDATE ml_learner_state
|
|
SET leader_token = %s,
|
|
lease_expires_at = NOW() + (%s * INTERVAL '1 second'),
|
|
heartbeat_at = NOW(),
|
|
run_started_at = NOW(),
|
|
run_deadline_at = NOW() + (%s * INTERVAL '1 second'),
|
|
last_terminal_reason = NULL,
|
|
data_version = %s,
|
|
updated_at = NOW()
|
|
WHERE singleton = TRUE
|
|
AND next_run_at <= NOW()
|
|
AND (
|
|
leader_token IS NULL
|
|
OR lease_expires_at IS NULL
|
|
OR lease_expires_at <= NOW()
|
|
)
|
|
RETURNING singleton,
|
|
next_training_at <= NOW() AS should_train,
|
|
model_version""",
|
|
[
|
|
token,
|
|
ML_LEASE_SECONDS,
|
|
ML_RUN_DEADLINE_SECONDS,
|
|
ML_DATA_VERSION,
|
|
],
|
|
)
|
|
row = cur.fetchone()
|
|
db.commit()
|
|
if row is None:
|
|
return None
|
|
return RunGuard(
|
|
token,
|
|
time.monotonic() + ML_RUN_DEADLINE_SECONDS,
|
|
should_train=bool(row['should_train']),
|
|
expected_model_version=str(row['model_version']),
|
|
)
|
|
|
|
|
|
@dataclass(frozen=True, order=True)
|
|
class ExtractionCursor:
|
|
observed_at: datetime
|
|
packet_hash: str
|
|
network: str
|
|
rx_node_id: str
|
|
topic: str
|
|
raw_hex: str
|
|
|
|
def sql_values(self) -> tuple:
|
|
return (
|
|
self.observed_at,
|
|
self.packet_hash,
|
|
self.network,
|
|
self.rx_node_id,
|
|
self.topic,
|
|
self.raw_hex,
|
|
)
|
|
|
|
|
|
def cursor_from_row(row) -> ExtractionCursor:
|
|
observed_at = row['observed_at']
|
|
if not isinstance(observed_at, datetime):
|
|
observed_at = datetime.fromisoformat(str(observed_at))
|
|
if observed_at.tzinfo is None:
|
|
observed_at = observed_at.replace(tzinfo=timezone.utc)
|
|
return ExtractionCursor(
|
|
observed_at=observed_at,
|
|
packet_hash=str(row['packet_hash'] or ''),
|
|
network=str(row['network'] or ''),
|
|
rx_node_id=str(row['rx_node_id'] or ''),
|
|
topic=str(row['topic'] or ''),
|
|
raw_hex=str(row['raw_hex'] or ''),
|
|
)
|
|
|
|
|
|
def get_checkpoint(db) -> ExtractionCursor:
|
|
with db.cursor() as cur:
|
|
cur.execute(
|
|
"""SELECT cursor_observed_at, cursor_packet_hash, cursor_network,
|
|
cursor_rx_node_id, cursor_topic, cursor_raw_hex
|
|
FROM ml_learner_state
|
|
WHERE singleton = TRUE"""
|
|
)
|
|
row = cur.fetchone()
|
|
if not row:
|
|
return ExtractionCursor(
|
|
datetime(1970, 1, 1, tzinfo=timezone.utc),
|
|
'',
|
|
'',
|
|
'',
|
|
'',
|
|
'',
|
|
)
|
|
return ExtractionCursor(
|
|
observed_at=row['cursor_observed_at'],
|
|
packet_hash=row['cursor_packet_hash'],
|
|
network=row['cursor_network'],
|
|
rx_node_id=row['cursor_rx_node_id'],
|
|
topic=row['cursor_topic'],
|
|
raw_hex=row['cursor_raw_hex'],
|
|
)
|
|
|
|
|
|
def set_checkpoint(db, cursor: ExtractionCursor, guard: RunGuard):
|
|
with db.cursor() as cur:
|
|
guard.fence_checkpoint(cur, cursor)
|
|
|
|
|
|
# ── Gold extraction ───────────────────────────────────────────────────────────
|
|
|
|
def _trim_terminal_hash(path_hashes: list[str], rx_node_id: str) -> list[str]:
|
|
"""Remove the observer's own hash that meshcore appends as the last entry."""
|
|
if not path_hashes:
|
|
return path_hashes
|
|
last = path_hashes[-1].upper()
|
|
if rx_node_id.upper().startswith(last):
|
|
return path_hashes[:-1]
|
|
return path_hashes
|
|
|
|
|
|
def extract_gold_paths(db, guard: RunGuard):
|
|
guard.check()
|
|
checkpoint = get_checkpoint(db)
|
|
log.info('Gold extraction from checkpoint %s', checkpoint)
|
|
|
|
with db.cursor() as cur:
|
|
cur.execute(
|
|
"""SELECT p.packet_hash, p.network, p.rx_node_id, p.topic, p.raw_hex,
|
|
p.path_hashes, p.path_hash_size_bytes, p.time AS observed_at,
|
|
receiver.iata AS rx_region
|
|
FROM packets p
|
|
LEFT JOIN nodes receiver ON receiver.node_id = p.rx_node_id
|
|
WHERE p.path_hash_size_bytes > 1
|
|
AND p.path_hashes IS NOT NULL
|
|
AND cardinality(p.path_hashes) > 1
|
|
AND ROW(
|
|
p.time,
|
|
p.packet_hash,
|
|
p.network,
|
|
COALESCE(p.rx_node_id, ''),
|
|
p.topic,
|
|
COALESCE(p.raw_hex, '')
|
|
) > ROW(%s, %s, %s, %s, %s, %s)
|
|
ORDER BY
|
|
p.time ASC,
|
|
p.packet_hash ASC,
|
|
p.network ASC,
|
|
COALESCE(p.rx_node_id, '') ASC,
|
|
p.topic ASC,
|
|
COALESCE(p.raw_hex, '') ASC
|
|
LIMIT %s""",
|
|
[*checkpoint.sql_values(), GOLD_BATCH],
|
|
)
|
|
rows = cur.fetchall()
|
|
|
|
guard.check()
|
|
if not rows:
|
|
log.info('Gold extraction: no new packets')
|
|
db.rollback()
|
|
return 0
|
|
|
|
# ── Collect all hashes to batch-resolve ──────────────────────────────────
|
|
# Map: (network scope, hash_upper) → [node row, ...]
|
|
hash_to_resolve: dict[str, set[str]] = defaultdict(set)
|
|
scope_networks: dict[str, set[str]] = defaultdict(set)
|
|
for row_index, row in enumerate(rows):
|
|
if row_index % 250 == 0:
|
|
guard.check()
|
|
scope = network_scope_key(row['network'])
|
|
scope_networks[scope].update(network_scope_values(row['network']))
|
|
hashes = [h.upper() for h in (row['path_hashes'] or [])]
|
|
hashes = _trim_terminal_hash(hashes, row['rx_node_id'])
|
|
hash_size = (row['path_hash_size_bytes'] or 1) * 2 # expected hex chars
|
|
for h in hashes:
|
|
if len(h) == hash_size:
|
|
hash_to_resolve[scope].add(h)
|
|
|
|
nodes_by_net_hash: dict[tuple[str, str], list[dict]] = {}
|
|
for scope, hashes in hash_to_resolve.items():
|
|
guard.check()
|
|
if not hashes:
|
|
continue
|
|
hash_list = sorted(hashes)
|
|
# Use a single query with array containment to match prefixes efficiently
|
|
conditions = ' OR '.join([f"upper(node_id) LIKE %s" for _ in hash_list])
|
|
with db.cursor() as cur:
|
|
cur.execute(
|
|
f"""SELECT node_id, network, lat, lon, elevation_m, last_seen, iata
|
|
FROM nodes
|
|
WHERE network = ANY(%s)
|
|
AND lat IS NOT NULL AND lon IS NOT NULL
|
|
AND lat != 0 AND lon != 0
|
|
AND ({conditions})""",
|
|
[list(scope_networks[scope])] + [h + '%' for h in hash_list],
|
|
)
|
|
node_rows = cur.fetchall()
|
|
|
|
for node in node_rows:
|
|
nid = node['node_id'].upper()
|
|
for h in hash_list:
|
|
if nid.startswith(h):
|
|
key = (scope, h)
|
|
if key not in nodes_by_net_hash:
|
|
nodes_by_net_hash[key] = []
|
|
nodes_by_net_hash[key].append(dict(node))
|
|
break
|
|
|
|
# ── Group observations by (packet_hash, network) ─────────────────────────
|
|
by_packet: dict[tuple[str, str], list[dict]] = defaultdict(list)
|
|
for row in rows:
|
|
by_packet[(row['packet_hash'], row['network'])].append(dict(row))
|
|
|
|
inserted = 0
|
|
|
|
for packet_index, ((packet_hash, network), observations) in enumerate(by_packet.items()):
|
|
if packet_index % 100 == 0:
|
|
guard.check()
|
|
# All observers for this packet
|
|
all_observer_ids = [o['rx_node_id'] for o in observations]
|
|
|
|
for obs in observations:
|
|
hashes = [h.upper() for h in (obs['path_hashes'] or [])]
|
|
hashes = _trim_terminal_hash(hashes, obs['rx_node_id'])
|
|
hash_size = (obs['path_hash_size_bytes'] or 1) * 2
|
|
|
|
# Validate hash lengths
|
|
valid_hashes = [h for h in hashes if len(h) == hash_size]
|
|
if len(valid_hashes) < 2:
|
|
continue
|
|
|
|
# Resolve each hop — require unique match
|
|
resolved: list[tuple[str, dict]] = [] # (hash, node)
|
|
ok = True
|
|
for h in valid_hashes:
|
|
candidates = nodes_by_net_hash.get((network_scope_key(network), h), [])
|
|
if len(candidates) != 1:
|
|
ok = False
|
|
break
|
|
resolved.append((h, candidates[0]))
|
|
|
|
if not ok or len(resolved) < 2:
|
|
continue
|
|
|
|
# Validate: no impossible adjacent hops
|
|
for i in range(len(resolved) - 1):
|
|
a = resolved[i][1]
|
|
b = resolved[i + 1][1]
|
|
d = dist_km(a['lat'], a['lon'], b['lat'], b['lon'])
|
|
if d > MAX_HOP_KM:
|
|
ok = False
|
|
break
|
|
|
|
if not ok:
|
|
continue
|
|
|
|
rx_region = obs.get('rx_region')
|
|
|
|
# Insert gold hop rows
|
|
for pos, (h, node) in enumerate(resolved):
|
|
node_id = node['node_id']
|
|
hash_2char = node_id.upper()[:2]
|
|
hash_4char = node_id.upper()[:4]
|
|
hash_6char = node_id.upper()[:6]
|
|
with db.cursor() as cur:
|
|
cur.execute(
|
|
"""INSERT INTO ml_gold_paths
|
|
(packet_hash, network, observed_at, hop_position,
|
|
true_node_id, hash_2char, hash_4char, hash_6char,
|
|
path_hash_size_bytes, observer_ids, rx_region)
|
|
VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s)
|
|
ON CONFLICT (packet_hash, hop_position, true_node_id) DO NOTHING""",
|
|
[packet_hash, network, obs['observed_at'], pos,
|
|
node_id, hash_2char, hash_4char, hash_6char,
|
|
obs['path_hash_size_bytes'], all_observer_ids, rx_region],
|
|
)
|
|
inserted += cur.rowcount
|
|
|
|
set_checkpoint(db, cursor_from_row(rows[-1]), guard)
|
|
guard.check()
|
|
db.commit()
|
|
log.info('Gold extraction: inserted %d new hop rows from %d packets', inserted, len(by_packet))
|
|
return len(rows)
|
|
|
|
|
|
# ── Feature building ──────────────────────────────────────────────────────────
|
|
|
|
def build_training_data(db, guard: RunGuard):
|
|
"""
|
|
Build (X, y, meta) arrays for training.
|
|
|
|
For each gold hop, generates one row per candidate node sharing the
|
|
same 1-byte (2-char) hash prefix. The true node gets label=1; all
|
|
others get label=0.
|
|
"""
|
|
# Select whole packets in SQL before transferring rows to the worker. The
|
|
# round-robin stratum rank retains rare network/prefix examples while the
|
|
# cumulative row budget places a hard bound on application memory.
|
|
guard.check()
|
|
with db.cursor() as cur:
|
|
cur.execute(
|
|
"""WITH packet_groups AS (
|
|
SELECT
|
|
network,
|
|
packet_hash,
|
|
MIN(hash_2char) AS primary_hash,
|
|
MAX(observed_at) AS latest_observed_at,
|
|
COUNT(*)::bigint AS gold_rows
|
|
FROM ml_gold_paths
|
|
GROUP BY network, packet_hash
|
|
),
|
|
stratified AS (
|
|
SELECT
|
|
packet_groups.*,
|
|
ROW_NUMBER() OVER (
|
|
PARTITION BY network, primary_hash
|
|
ORDER BY latest_observed_at DESC, packet_hash
|
|
) AS stratum_rank
|
|
FROM packet_groups
|
|
),
|
|
budgeted AS (
|
|
SELECT
|
|
stratified.*,
|
|
SUM(gold_rows) OVER (
|
|
ORDER BY stratum_rank, latest_observed_at DESC, network, packet_hash
|
|
) AS running_rows
|
|
FROM stratified
|
|
),
|
|
selected_packets AS (
|
|
SELECT network, packet_hash
|
|
FROM budgeted
|
|
WHERE running_rows <= %s
|
|
)
|
|
SELECT g.id, g.packet_hash, g.network, g.observed_at,
|
|
g.hop_position, g.true_node_id, g.hash_2char,
|
|
g.path_hash_size_bytes,
|
|
COUNT(*) OVER (PARTITION BY g.packet_hash, g.network) as path_length
|
|
FROM ml_gold_paths g
|
|
JOIN selected_packets selected
|
|
ON selected.network = g.network
|
|
AND selected.packet_hash = g.packet_hash
|
|
ORDER BY g.network, g.hash_2char, g.observed_at, g.id""",
|
|
[MAX_TRAINING_GOLD_ROWS],
|
|
)
|
|
gold_rows = cur.fetchall()
|
|
guard.check()
|
|
log.info('SQL-bounded training sample retained %d gold rows', len(gold_rows))
|
|
|
|
if len(gold_rows) < MIN_GOLD_ROWS:
|
|
log.info('Only %d gold rows, skipping training', len(gold_rows))
|
|
return None, None, None, None
|
|
|
|
# Load all candidate nodes per combined network scope + 1-byte prefix.
|
|
hash_pairs: dict[tuple[str, str], set[str]] = defaultdict(set)
|
|
for row in gold_rows:
|
|
scope = network_scope_key(row['network'])
|
|
hash_pairs[(scope, row['hash_2char'])].update(network_scope_values(row['network']))
|
|
|
|
candidates_map: dict[tuple[str, str], list[dict]] = {}
|
|
for (scope, hash_2char), networks in hash_pairs.items():
|
|
guard.check()
|
|
with db.cursor() as cur:
|
|
cur.execute(
|
|
"""SELECT node_id, network AS node_network, elevation_m, last_seen
|
|
FROM nodes
|
|
WHERE network = ANY(%s)
|
|
AND upper(node_id) LIKE %s
|
|
AND lat IS NOT NULL AND lon IS NOT NULL""",
|
|
[list(networks), hash_2char + '%'],
|
|
)
|
|
candidates_map[(scope, hash_2char)] = cur.fetchall()
|
|
|
|
now = datetime.now(timezone.utc)
|
|
X_rows = []
|
|
y_rows = []
|
|
# (packet_network, packet_hash, hop_position, hash_2char, candidate_node_id,
|
|
# true_node_id, gold_id, path_length, candidate_network)
|
|
meta_rows = []
|
|
|
|
for row_index, r in enumerate(gold_rows):
|
|
if row_index % 500 == 0:
|
|
guard.check()
|
|
rid = r['id']
|
|
network = r['network']
|
|
hash_2char = r['hash_2char']
|
|
true_node_id = r['true_node_id']
|
|
candidates = candidates_map.get((network_scope_key(network), hash_2char), [])
|
|
if not candidates:
|
|
continue
|
|
|
|
collision_count = len(candidates)
|
|
|
|
for cand in candidates:
|
|
cand_id = cand['node_id']
|
|
label = 1 if cand_id == true_node_id else 0
|
|
|
|
# Days since last seen
|
|
ls = cand['last_seen']
|
|
if ls and hasattr(ls, 'tzinfo'):
|
|
if ls.tzinfo is None:
|
|
ls = ls.replace(tzinfo=timezone.utc)
|
|
days_since = (now - ls).total_seconds() / 86400.0
|
|
else:
|
|
days_since = 999.0
|
|
|
|
feat = [
|
|
collision_count,
|
|
float(cand['elevation_m'] or 0),
|
|
min(days_since, 999.0),
|
|
1 if days_since < 7 else 0,
|
|
int(r['hop_position']),
|
|
int(r['path_length'] or 1),
|
|
1, # simulated 1-byte path hash size
|
|
]
|
|
X_rows.append(feat)
|
|
y_rows.append(label)
|
|
meta_rows.append((
|
|
network,
|
|
r['packet_hash'],
|
|
int(r['hop_position']),
|
|
hash_2char,
|
|
cand_id,
|
|
true_node_id,
|
|
rid,
|
|
int(r['path_length'] or 1),
|
|
cand['node_network'],
|
|
))
|
|
|
|
if not X_rows:
|
|
return None, None, None, None
|
|
|
|
X = np.array(X_rows, dtype=np.float32)
|
|
y = np.array(y_rows, dtype=np.int32)
|
|
# gold_ids: parallel array of ml_gold_paths.id, used for grouping in evaluation
|
|
gold_ids = np.array([m[6] for m in meta_rows], dtype=np.int64)
|
|
return X, y, meta_rows, gold_ids
|
|
|
|
|
|
FEATURE_NAMES = [
|
|
'collision_count', 'elevation_m', 'days_since_seen', 'is_online_recent',
|
|
'hop_position', 'path_length', 'simulated_path_hash_size_bytes',
|
|
]
|
|
|
|
|
|
# ── Training ──────────────────────────────────────────────────────────────────
|
|
|
|
def train_variant(
|
|
X_train: np.ndarray, y_train: np.ndarray,
|
|
X_val: np.ndarray, y_val: np.ndarray,
|
|
params: dict,
|
|
) -> object | None:
|
|
"""Train one LightGBM variant with given hyperparams and return calibrated model."""
|
|
if len(set(y_train)) < 2 or len(set(y_val)) < 2:
|
|
return None
|
|
|
|
rng = np.random.default_rng(RANDOM_SEED)
|
|
train_idx = np.arange(len(y_train))
|
|
rng.shuffle(train_idx)
|
|
split = int(len(train_idx) * 0.8)
|
|
split = max(1, min(len(train_idx) - 1, split))
|
|
fit_idx = train_idx[:split]
|
|
cal_idx = train_idx[split:]
|
|
|
|
X_fit, y_fit = X_train[fit_idx], y_train[fit_idx]
|
|
X_cal, y_cal = X_train[cal_idx], y_train[cal_idx]
|
|
if len(set(y_fit.tolist())) < 2:
|
|
X_fit, y_fit = X_train, y_train
|
|
|
|
base = LGBMClassifier(
|
|
objective='binary',
|
|
n_jobs=2,
|
|
is_unbalance=True,
|
|
verbose=-1,
|
|
**{k: v for k, v in params.items() if k != 'bagging_freq' or v > 0},
|
|
)
|
|
try:
|
|
eval_set = [(X_cal, y_cal)] if len(set(y_cal.tolist())) >= 2 else None
|
|
base.fit(X_fit, y_fit, eval_set=eval_set, callbacks=[])
|
|
except Exception as e:
|
|
log.warning('Variant training error: %s', e)
|
|
return None
|
|
|
|
if len(set(y_cal.tolist())) < 2:
|
|
return base
|
|
|
|
try:
|
|
model = CalibratedClassifierCV(FrozenEstimator(base), method='isotonic')
|
|
model.fit(X_cal, y_cal)
|
|
return model
|
|
except (ValueError, RuntimeError) as exc:
|
|
log.warning('Calibration failed; using uncalibrated variant: %s', exc)
|
|
return base
|
|
|
|
|
|
def train_final_variant(X: np.ndarray, y: np.ndarray, params: dict) -> object | None:
|
|
"""Train a final LightGBM model on all gold rows without calibration splits."""
|
|
if len(set(y.tolist())) < 2:
|
|
return None
|
|
|
|
model = LGBMClassifier(
|
|
objective='binary',
|
|
n_jobs=2,
|
|
is_unbalance=True,
|
|
verbose=-1,
|
|
**{k: v for k, v in params.items() if k != 'bagging_freq' or v > 0},
|
|
)
|
|
try:
|
|
model.fit(X, y)
|
|
return model
|
|
except Exception as e:
|
|
log.warning('Final variant training error: %s', e)
|
|
return None
|
|
|
|
|
|
def collect_path_predictions(model, X: np.ndarray, y: np.ndarray, meta_rows: list) -> tuple[list[dict], dict]:
|
|
"""
|
|
Run the model over candidate rows and aggregate predictions back into
|
|
packet paths.
|
|
|
|
A candidate row is still the model's unit of inference, but a packet path is
|
|
the unit of evaluation. A packet is only complete when every expected hop
|
|
is present and the top-ranked candidate for every hop is the true node.
|
|
"""
|
|
probs = model.predict_proba(X)[:, 1]
|
|
|
|
hop_groups: dict[int, list[int]] = defaultdict(list)
|
|
packets: dict[tuple[str, str], dict] = {}
|
|
|
|
for i, meta in enumerate(meta_rows):
|
|
net, packet_hash, hop_pos, h2, cid, true_id, gid, path_len, _cand_net = meta
|
|
hop_groups[int(gid)].append(i)
|
|
packet_key = (net, packet_hash)
|
|
if packet_key not in packets:
|
|
packets[packet_key] = {
|
|
'network': net,
|
|
'packet_hash': packet_hash,
|
|
'expected_hops': int(path_len),
|
|
'predicted_hops': {},
|
|
}
|
|
else:
|
|
packets[packet_key]['expected_hops'] = max(
|
|
int(packets[packet_key]['expected_hops']),
|
|
int(path_len),
|
|
)
|
|
|
|
predictions: list[dict] = []
|
|
|
|
for gid, idxs in hop_groups.items():
|
|
if not idxs:
|
|
continue
|
|
idx_arr = np.array(idxs, dtype=np.int64)
|
|
labels = y[idx_arr]
|
|
if labels.sum() != 1:
|
|
continue
|
|
|
|
group_probs = probs[idx_arr]
|
|
best_idx = int(idx_arr[int(np.argmax(group_probs))])
|
|
top3_local = np.argsort(group_probs)[::-1][:3]
|
|
top3_correct = bool(labels[top3_local].sum() > 0)
|
|
|
|
net, packet_hash, hop_pos, h2, cid, true_id, _gid, path_len, _cand_net = meta_rows[best_idx]
|
|
packet_key = (net, packet_hash)
|
|
correct = bool(y[best_idx] == 1)
|
|
|
|
packets[packet_key]['predicted_hops'][int(hop_pos)] = correct
|
|
predictions.append({
|
|
'network': net,
|
|
'packet_hash': packet_hash,
|
|
'packet_key': packet_key,
|
|
'hop_position': int(hop_pos),
|
|
'hash_2char': h2,
|
|
'node_id': cid,
|
|
'true_node_id': true_id,
|
|
'gold_id': gid,
|
|
'correct': correct,
|
|
'top3_correct': top3_correct,
|
|
'probability': float(probs[best_idx]),
|
|
})
|
|
|
|
for packet in packets.values():
|
|
expected = max(1, int(packet['expected_hops']))
|
|
predicted_hops: dict[int, bool] = packet['predicted_hops']
|
|
correct_hops = sum(1 for pos in range(expected) if predicted_hops.get(pos, False))
|
|
predicted_count = len(predicted_hops)
|
|
packet['correct_hops'] = correct_hops
|
|
packet['predicted_hops_count'] = predicted_count
|
|
packet['complete'] = predicted_count >= expected and correct_hops == expected
|
|
packet['completion'] = correct_hops / expected
|
|
|
|
return predictions, packets
|
|
|
|
|
|
def evaluate_path_metrics(model, X: np.ndarray, y: np.ndarray, meta_rows: list) -> dict:
|
|
return evaluate_path_details(model, X, y, meta_rows)[0]
|
|
|
|
|
|
def evaluate_path_details(model, X: np.ndarray, y: np.ndarray, meta_rows: list) -> tuple[dict, list[dict], dict]:
|
|
predictions, packets = collect_path_predictions(model, X, y, meta_rows)
|
|
hop_total = len(predictions)
|
|
hop_correct = sum(1 for p in predictions if p['correct'])
|
|
top3_correct = sum(1 for p in predictions if p['top3_correct'])
|
|
|
|
packet_total = len(packets)
|
|
complete_paths = sum(1 for p in packets.values() if p['complete'])
|
|
total_expected_hops = sum(int(p['expected_hops']) for p in packets.values())
|
|
total_correct_path_hops = sum(int(p['correct_hops']) for p in packets.values())
|
|
|
|
metrics = {
|
|
'hop_total': hop_total,
|
|
'hop_correct': hop_correct,
|
|
'hop_accuracy': hop_correct / hop_total if hop_total else 0.0,
|
|
'hop_top3_accuracy': top3_correct / hop_total if hop_total else 0.0,
|
|
'packet_total': packet_total,
|
|
'complete_paths': complete_paths,
|
|
'complete_path_accuracy': complete_paths / packet_total if packet_total else 0.0,
|
|
'mean_path_completion': (
|
|
total_correct_path_hops / total_expected_hops
|
|
if total_expected_hops else 0.0
|
|
),
|
|
}
|
|
return metrics, predictions, packets
|
|
|
|
|
|
def persist_variant_evaluation(
|
|
db,
|
|
training_run_id: str,
|
|
generation: int,
|
|
variant_rank: int,
|
|
params: dict,
|
|
all_metrics: dict,
|
|
val_metrics: dict,
|
|
packets: dict,
|
|
guard: RunGuard,
|
|
):
|
|
guard.check()
|
|
packet_rows = [
|
|
(
|
|
training_run_id,
|
|
GLOBAL_NETWORK,
|
|
generation,
|
|
variant_rank,
|
|
packet['network'],
|
|
packet['packet_hash'],
|
|
int(packet['expected_hops']),
|
|
int(packet['predicted_hops_count']),
|
|
int(packet['correct_hops']),
|
|
bool(packet['complete']),
|
|
float(packet['completion']),
|
|
)
|
|
for packet in packets.values()
|
|
]
|
|
|
|
with db.cursor() as cur:
|
|
cur.execute(
|
|
"""INSERT INTO ml_model_variant_runs
|
|
(training_run_id, model_network, generation, variant_rank,
|
|
population_size, hyperparams, evaluated_packets,
|
|
evaluated_hops, hop_accuracy, hop_top3_accuracy,
|
|
complete_path_accuracy, mean_path_completion,
|
|
val_evaluated_packets, val_evaluated_hops,
|
|
val_hop_accuracy, val_hop_top3_accuracy,
|
|
val_complete_path_accuracy, val_mean_path_completion,
|
|
created_at)
|
|
VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s,
|
|
%s, %s, %s, %s, %s, %s, NOW())
|
|
ON CONFLICT (training_run_id, variant_rank) DO UPDATE SET
|
|
hyperparams = EXCLUDED.hyperparams,
|
|
evaluated_packets = EXCLUDED.evaluated_packets,
|
|
evaluated_hops = EXCLUDED.evaluated_hops,
|
|
hop_accuracy = EXCLUDED.hop_accuracy,
|
|
hop_top3_accuracy = EXCLUDED.hop_top3_accuracy,
|
|
complete_path_accuracy = EXCLUDED.complete_path_accuracy,
|
|
mean_path_completion = EXCLUDED.mean_path_completion,
|
|
val_evaluated_packets = EXCLUDED.val_evaluated_packets,
|
|
val_evaluated_hops = EXCLUDED.val_evaluated_hops,
|
|
val_hop_accuracy = EXCLUDED.val_hop_accuracy,
|
|
val_hop_top3_accuracy = EXCLUDED.val_hop_top3_accuracy,
|
|
val_complete_path_accuracy = EXCLUDED.val_complete_path_accuracy,
|
|
val_mean_path_completion = EXCLUDED.val_mean_path_completion""",
|
|
[
|
|
training_run_id,
|
|
GLOBAL_NETWORK,
|
|
generation,
|
|
variant_rank,
|
|
POPULATION_SIZE,
|
|
json.dumps(params),
|
|
all_metrics['packet_total'],
|
|
all_metrics['hop_total'],
|
|
all_metrics['hop_accuracy'],
|
|
all_metrics['hop_top3_accuracy'],
|
|
all_metrics['complete_path_accuracy'],
|
|
all_metrics['mean_path_completion'],
|
|
val_metrics['packet_total'],
|
|
val_metrics['hop_total'],
|
|
val_metrics['hop_accuracy'],
|
|
val_metrics['hop_top3_accuracy'],
|
|
val_metrics['complete_path_accuracy'],
|
|
val_metrics['mean_path_completion'],
|
|
],
|
|
)
|
|
|
|
if packet_rows:
|
|
psycopg2.extras.execute_values(
|
|
cur,
|
|
"""INSERT INTO ml_model_variant_packet_results
|
|
(training_run_id, model_network, generation, variant_rank,
|
|
packet_network, packet_hash, expected_hops,
|
|
predicted_hops, correct_hops, complete_path,
|
|
path_completion)
|
|
VALUES %s
|
|
ON CONFLICT (training_run_id, variant_rank, packet_network, packet_hash)
|
|
DO UPDATE SET
|
|
expected_hops = EXCLUDED.expected_hops,
|
|
predicted_hops = EXCLUDED.predicted_hops,
|
|
correct_hops = EXCLUDED.correct_hops,
|
|
complete_path = EXCLUDED.complete_path,
|
|
path_completion = EXCLUDED.path_completion""",
|
|
packet_rows,
|
|
page_size=1000,
|
|
)
|
|
guard.check()
|
|
db.commit()
|
|
|
|
|
|
def split_by_packet(meta_rows: list, train_fraction: float = 0.8) -> tuple[np.ndarray, np.ndarray]:
|
|
"""Split candidate rows by packet so a full path is entirely train or val."""
|
|
packet_keys = sorted({(m[0], m[1]) for m in meta_rows})
|
|
if len(packet_keys) < 2:
|
|
empty = np.zeros(len(meta_rows), dtype=bool)
|
|
return ~empty, empty
|
|
|
|
rng = random.Random(RANDOM_SEED)
|
|
rng.shuffle(packet_keys)
|
|
split = int(len(packet_keys) * train_fraction)
|
|
split = max(1, min(len(packet_keys) - 1, split))
|
|
|
|
train_packets = set(packet_keys[:split])
|
|
train_mask = np.array([(m[0], m[1]) in train_packets for m in meta_rows], dtype=bool)
|
|
val_mask = ~train_mask
|
|
return train_mask, val_mask
|
|
|
|
|
|
# ── Promotion ─────────────────────────────────────────────────────────────────
|
|
|
|
def get_current_best_accuracy(db) -> float:
|
|
with db.cursor() as cur:
|
|
cur.execute(
|
|
"""SELECT complete_path_accuracy AS champion_score
|
|
FROM ml_model_versions
|
|
WHERE network = %s AND is_active = TRUE ORDER BY promoted_at DESC LIMIT 1""",
|
|
[GLOBAL_NETWORK],
|
|
)
|
|
row = cur.fetchone()
|
|
return float(row['champion_score']) if row and row['champion_score'] is not None else 0.0
|
|
|
|
|
|
def evaluate_current_champion(db, X: np.ndarray, y: np.ndarray, meta_rows: list) -> dict | None:
|
|
"""Replay the active champion on the current gold corpus."""
|
|
with db.cursor() as cur:
|
|
cur.execute(
|
|
"""SELECT version, generation, variant_rank, model_artifact
|
|
FROM ml_model_versions
|
|
WHERE network = %s AND is_active = TRUE
|
|
ORDER BY promoted_at DESC LIMIT 1""",
|
|
[GLOBAL_NETWORK],
|
|
)
|
|
row = cur.fetchone()
|
|
|
|
if not row or not row['model_artifact']:
|
|
return None
|
|
|
|
try:
|
|
model = joblib.load(io.BytesIO(bytes(row['model_artifact'])))
|
|
metrics = evaluate_path_metrics(model, X, y, meta_rows)
|
|
metrics['version'] = row['version']
|
|
metrics['generation'] = int(row['generation'])
|
|
metrics['variant_rank'] = int(row['variant_rank'])
|
|
return metrics
|
|
except Exception as e:
|
|
log.warning('Current champion replay failed: %s', e, exc_info=True)
|
|
return None
|
|
|
|
|
|
def promotion_score(metrics: dict) -> tuple[float, float, float]:
|
|
return (
|
|
float(metrics['complete_path_accuracy']),
|
|
float(metrics['hop_accuracy']),
|
|
float(metrics['mean_path_completion']),
|
|
)
|
|
|
|
|
|
def score_beats(candidate: tuple[float, float, float],
|
|
incumbent: tuple[float, float, float],
|
|
min_delta: float = 0.0) -> bool:
|
|
if candidate[0] > incumbent[0] + min_delta:
|
|
return True
|
|
if candidate[0] + min_delta < incumbent[0]:
|
|
return False
|
|
return candidate[1:] > incumbent[1:]
|
|
|
|
|
|
def promote_model(db, model, val_metrics: dict, all_metrics: dict, gold_count: int,
|
|
X: np.ndarray, y: np.ndarray, meta_rows: list,
|
|
gold_ids: np.ndarray,
|
|
hyperparams: dict | None = None,
|
|
generation: int = 1, variant_rank: int = 1,
|
|
guard: RunGuard | None = None):
|
|
if guard is None:
|
|
raise RuntimeError('a leader run guard is required for model publication')
|
|
guard.check()
|
|
version = (
|
|
datetime.now(timezone.utc).strftime('%Y%m%d_%H%M%S_%f')
|
|
+ f'_global_{guard.token[:8]}'
|
|
)
|
|
|
|
# Serialize model
|
|
buf = io.BytesIO()
|
|
joblib.dump(model, buf)
|
|
artifact = buf.getvalue()
|
|
|
|
with db.cursor() as cur:
|
|
cur.execute("UPDATE ml_model_versions SET is_active = FALSE WHERE is_active = TRUE")
|
|
cur.execute(
|
|
"""INSERT INTO ml_model_versions
|
|
(version, network, trained_at, gold_paths_used, top1_accuracy,
|
|
top3_accuracy, is_active, promoted_at, model_artifact,
|
|
hyperparams, generation, variant_rank, population_size,
|
|
evaluated_packets, evaluated_hops, complete_path_accuracy,
|
|
mean_path_completion, data_version)
|
|
VALUES (%s, %s, NOW(), %s, %s, %s, TRUE, NOW(), %s, %s, %s, %s, %s,
|
|
%s, %s, %s, %s, %s)""",
|
|
[
|
|
version,
|
|
GLOBAL_NETWORK,
|
|
gold_count,
|
|
all_metrics['hop_accuracy'],
|
|
all_metrics['hop_top3_accuracy'],
|
|
psycopg2.Binary(artifact),
|
|
json.dumps(hyperparams or {}),
|
|
generation,
|
|
variant_rank,
|
|
POPULATION_SIZE,
|
|
all_metrics['packet_total'],
|
|
all_metrics['hop_total'],
|
|
all_metrics['complete_path_accuracy'],
|
|
all_metrics['mean_path_completion'],
|
|
ML_DATA_VERSION,
|
|
],
|
|
)
|
|
|
|
probs = model.predict_proba(X)[:, 1]
|
|
predictions, packets = collect_path_predictions(model, X, y, meta_rows)
|
|
selected_predictions = {
|
|
(pred['gold_id'], pred['node_id']): pred
|
|
for pred in predictions
|
|
}
|
|
score_rows: dict[tuple[str, str, str], dict] = defaultdict(
|
|
lambda: {
|
|
'observed': 0,
|
|
'correct': 0,
|
|
'prob_sum': 0.0,
|
|
'selected': 0,
|
|
'selected_correct': 0,
|
|
'packets': set(),
|
|
'complete_packets': set(),
|
|
}
|
|
)
|
|
for i, meta in enumerate(meta_rows):
|
|
packet_net, packet_hash, _hop_pos, h2, cid, _true_id, gid, _path_len, cand_net = meta
|
|
score_net = cand_net or packet_net
|
|
packet_key = (packet_net, packet_hash)
|
|
key = (score_net, h2, cid)
|
|
row = score_rows[key]
|
|
row['observed'] += 1
|
|
row['prob_sum'] += float(probs[i])
|
|
if y[i] == 1:
|
|
row['correct'] += 1
|
|
row['packets'].add(packet_key)
|
|
|
|
selected = selected_predictions.get((int(gid), cid))
|
|
if selected:
|
|
row['selected'] += 1
|
|
if selected['correct']:
|
|
row['selected_correct'] += 1
|
|
if packets[packet_key]['complete']:
|
|
row['complete_packets'].add(packet_key)
|
|
|
|
scores_written = 0
|
|
with db.cursor() as cur:
|
|
cur.execute("DELETE FROM ml_path_prefix_scores WHERE model_version != %s", [version])
|
|
for (net, h2, node_id), counts in score_rows.items():
|
|
obs = int(counts['observed'])
|
|
correct = int(counts['correct'])
|
|
avg_model_prob = float(counts['prob_sum']) / obs if obs else 0.0
|
|
# Persist a conservative score for every candidate mapping the
|
|
# model considered. The score cannot exceed the empirical rate at
|
|
# which this node was actually correct for the 1-byte prefix.
|
|
empirical_correct_rate = (correct + 1) / (obs + 2)
|
|
score = min(avg_model_prob, empirical_correct_rate)
|
|
if obs < MIN_OBSERVATION_COUNT or score < CONFIDENCE_THRESHOLD:
|
|
continue
|
|
cur.execute(
|
|
"""INSERT INTO ml_path_prefix_scores
|
|
(network, hash_2char, node_id, score, observation_count,
|
|
correct_count, packet_count, complete_path_count,
|
|
model_version, updated_at)
|
|
VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, NOW())
|
|
ON CONFLICT (network, hash_2char, node_id) DO UPDATE
|
|
SET score = EXCLUDED.score,
|
|
observation_count = EXCLUDED.observation_count,
|
|
correct_count = EXCLUDED.correct_count,
|
|
packet_count = EXCLUDED.packet_count,
|
|
complete_path_count = EXCLUDED.complete_path_count,
|
|
model_version = EXCLUDED.model_version,
|
|
updated_at = NOW()""",
|
|
[
|
|
net,
|
|
h2,
|
|
node_id,
|
|
score,
|
|
obs,
|
|
correct,
|
|
len(counts['packets']),
|
|
len(counts['complete_packets']),
|
|
version,
|
|
],
|
|
)
|
|
scores_written += 1
|
|
|
|
with db.cursor() as cur:
|
|
guard.fence_publication(cur, version)
|
|
db.commit()
|
|
log.info(
|
|
'Champion promoted network=%s generation=%d variant=%d/%d val_hop=%.3f val_top3=%.3f val_full_path=%.3f val_path_completion=%.3f all_hop=%.3f all_full_path=%.3f scores_written=%d confidence_threshold=%.2f params=%s',
|
|
GLOBAL_NETWORK, generation, variant_rank, POPULATION_SIZE,
|
|
val_metrics['hop_accuracy'], val_metrics['hop_top3_accuracy'],
|
|
val_metrics['complete_path_accuracy'], val_metrics['mean_path_completion'],
|
|
all_metrics['hop_accuracy'], all_metrics['complete_path_accuracy'],
|
|
scores_written, CONFIDENCE_THRESHOLD, json.dumps(hyperparams or {}),
|
|
)
|
|
|
|
|
|
def cleanup_training_artifacts(db, guard: RunGuard):
|
|
"""Keep recent diagnostics and active model artifacts; trim old bulk data."""
|
|
deleted_packet_results = 0
|
|
cleared_model_artifacts = 0
|
|
cleaned_generations: list[int] = []
|
|
|
|
guard.check()
|
|
with db.cursor() as cur:
|
|
if RETAIN_VARIANT_RESULT_GENERATIONS > 0:
|
|
cur.execute(
|
|
"""
|
|
WITH bounds AS (
|
|
SELECT GREATEST(0, COALESCE(MAX(generation), 0) - %s + 1) AS min_generation
|
|
FROM ml_model_variant_runs
|
|
),
|
|
active_generations AS (
|
|
SELECT DISTINCT generation
|
|
FROM ml_model_versions
|
|
WHERE is_active = TRUE
|
|
)
|
|
SELECT r.generation
|
|
FROM ml_model_variant_runs r, bounds b
|
|
WHERE b.min_generation > 0
|
|
AND r.generation < b.min_generation
|
|
AND NOT EXISTS (
|
|
SELECT 1
|
|
FROM active_generations a
|
|
WHERE a.generation = r.generation
|
|
)
|
|
GROUP BY r.generation
|
|
ORDER BY r.generation
|
|
LIMIT %s
|
|
""",
|
|
[RETAIN_VARIANT_RESULT_GENERATIONS, max(1, CLEANUP_GENERATION_BATCH_SIZE)],
|
|
)
|
|
cleaned_generations = [int(row['generation']) for row in cur.fetchall()]
|
|
|
|
if cleaned_generations:
|
|
cur.execute(
|
|
"""
|
|
DELETE FROM ml_model_variant_packet_results
|
|
WHERE generation = ANY(%s)
|
|
""",
|
|
[cleaned_generations],
|
|
)
|
|
deleted_packet_results = cur.rowcount
|
|
|
|
if RETAIN_MODEL_ARTIFACT_GENERATIONS > 0:
|
|
cur.execute(
|
|
"""
|
|
WITH bounds AS (
|
|
SELECT GREATEST(0, COALESCE(MAX(generation), 0) - %s + 1) AS min_generation
|
|
FROM ml_model_versions
|
|
)
|
|
UPDATE ml_model_versions m
|
|
SET model_artifact = NULL
|
|
FROM bounds b
|
|
WHERE b.min_generation > 0
|
|
AND m.generation < b.min_generation
|
|
AND m.is_active = FALSE
|
|
AND m.model_artifact IS NOT NULL
|
|
""",
|
|
[RETAIN_MODEL_ARTIFACT_GENERATIONS],
|
|
)
|
|
cleared_model_artifacts = cur.rowcount
|
|
|
|
guard.check()
|
|
db.commit()
|
|
|
|
if deleted_packet_results or cleared_model_artifacts:
|
|
log.info(
|
|
'Cleaned ML training artifacts: deleted_packet_results=%d cleared_model_artifacts=%d cleaned_generations=%s retain_variant_generations=%d retain_artifact_generations=%d',
|
|
deleted_packet_results, cleared_model_artifacts,
|
|
cleaned_generations,
|
|
RETAIN_VARIANT_RESULT_GENERATIONS, RETAIN_MODEL_ARTIFACT_GENERATIONS,
|
|
)
|
|
|
|
|
|
# ── Main loop ─────────────────────────────────────────────────────────────────
|
|
|
|
def run_training_cycle(db, guard: RunGuard):
|
|
guard.check()
|
|
log.info('Starting training cycle')
|
|
result = build_training_data(db, guard)
|
|
if result[0] is None:
|
|
log.info('Insufficient training data, skipping')
|
|
return
|
|
X, y, meta_rows, gold_ids = result
|
|
|
|
if len(y) < MIN_GOLD_ROWS:
|
|
log.info('Only %d rows, skipping', len(y))
|
|
return
|
|
|
|
# Train one global model across all networks combined.
|
|
# Split by packet so whole paths never straddle train/val.
|
|
n = len(y)
|
|
train_mask, val_mask = split_by_packet(meta_rows)
|
|
X_train, X_val = X[train_mask], X[val_mask]
|
|
y_train, y_val = y[train_mask], y[val_mask]
|
|
gids_train, gids_val = gold_ids[train_mask], gold_ids[val_mask]
|
|
meta_train = [m for m, keep in zip(meta_rows, train_mask) if keep]
|
|
meta_val = [m for m, keep in zip(meta_rows, val_mask) if keep]
|
|
|
|
if len(y_train) == 0 or len(y_val) == 0:
|
|
log.warning('Train/validation split produced an empty side, skipping')
|
|
return
|
|
|
|
if len(set(y_train.tolist())) < 2 or len(set(y_val.tolist())) < 2:
|
|
log.warning('Class diversity too low globally, skipping')
|
|
return
|
|
|
|
groups_val: dict[int, list[int]] = defaultdict(list)
|
|
for i, gid in enumerate(gids_val):
|
|
groups_val[int(gid)].append(i)
|
|
ambiguous_count = sum(1 for idxs in groups_val.values() if len(idxs) > 1)
|
|
train_packets = {(m[0], m[1]) for m in meta_train}
|
|
val_packets = {(m[0], m[1]) for m in meta_val}
|
|
all_packets = {(m[0], m[1]) for m in meta_rows}
|
|
|
|
generation = get_current_generation(db) + 1
|
|
training_run_id = datetime.now(timezone.utc).strftime('%Y%m%d_%H%M%S') + f'_gen{generation}'
|
|
champion_params = get_champion_params(db)
|
|
population = create_population(champion_params, generation)
|
|
stored_current_best = get_current_best_accuracy(db)
|
|
champion_current_metrics = evaluate_current_champion(db, X, y, meta_rows)
|
|
champion_val_metrics = evaluate_current_champion(db, X_val, y_val, meta_val)
|
|
current_all_best = (
|
|
champion_current_metrics['complete_path_accuracy']
|
|
if champion_current_metrics
|
|
else stored_current_best
|
|
)
|
|
current_val_best = (
|
|
champion_val_metrics['complete_path_accuracy']
|
|
if champion_val_metrics
|
|
else 0.0
|
|
)
|
|
|
|
networks_repr = ', '.join(sorted({m[0] for m in meta_rows}))
|
|
log.info(
|
|
'Generation %d network=%s population=%d rows=%d train_rows=%d val_rows=%d packets=%d train_packets=%d val_packets=%d gold_hops=%d train_hops=%d val_hops=%d positives=%d ambiguous_hops=%d champion_val_full_path=%.3f champion_all_full_path=%.3f stored_best_full_path=%.3f networks=[%s]',
|
|
generation, GLOBAL_NETWORK, POPULATION_SIZE, n, len(y_train), len(y_val),
|
|
len(all_packets), len(train_packets), len(val_packets),
|
|
len(set(gold_ids.tolist())), len(set(gids_train.tolist())), len(set(gids_val.tolist())),
|
|
int(sum(y)), ambiguous_count, current_val_best, current_all_best,
|
|
stored_current_best, networks_repr,
|
|
)
|
|
if champion_current_metrics:
|
|
log.info(
|
|
'Active champion replay on current corpus version=%s generation=%d variant=%d packets=%d hop=%.3f full_path=%.3f path_completion=%.3f',
|
|
champion_current_metrics['version'],
|
|
champion_current_metrics['generation'],
|
|
champion_current_metrics['variant_rank'],
|
|
champion_current_metrics['packet_total'],
|
|
champion_current_metrics['hop_accuracy'],
|
|
champion_current_metrics['complete_path_accuracy'],
|
|
champion_current_metrics['mean_path_completion'],
|
|
)
|
|
if champion_val_metrics:
|
|
log.info(
|
|
'Active champion replay on current validation split version=%s generation=%d variant=%d packets=%d hop=%.3f full_path=%.3f path_completion=%.3f',
|
|
champion_val_metrics['version'],
|
|
champion_val_metrics['generation'],
|
|
champion_val_metrics['variant_rank'],
|
|
champion_val_metrics['packet_total'],
|
|
champion_val_metrics['hop_accuracy'],
|
|
champion_val_metrics['complete_path_accuracy'],
|
|
champion_val_metrics['mean_path_completion'],
|
|
)
|
|
|
|
best_model = None
|
|
best_val_metrics: dict | None = None
|
|
best_all_metrics: dict | None = None
|
|
best_params: dict = champion_params
|
|
best_rank = 0
|
|
best_score = (-1.0, -1.0, -1.0, -1.0)
|
|
|
|
for rank, params in enumerate(population, start=1):
|
|
guard.check()
|
|
model = train_variant(X_train, y_train, X_val, y_val, params)
|
|
if model is None:
|
|
continue
|
|
|
|
val_metrics = evaluate_path_metrics(model, X_val, y_val, meta_val)
|
|
all_metrics, _all_predictions, all_packets_detail = evaluate_path_details(model, X, y, meta_rows)
|
|
persist_variant_evaluation(
|
|
db, training_run_id, generation, rank, params,
|
|
all_metrics, val_metrics, all_packets_detail, guard,
|
|
)
|
|
|
|
log.info(
|
|
'Generation %d variant %d/%d network=%s val_hop=%d/%d %.3f val_top3=%.3f val_full_path=%d/%d %.3f val_path_completion=%.3f all_hop=%d/%d %.3f all_full_path=%d/%d %.3f params=%s',
|
|
generation, rank, POPULATION_SIZE, GLOBAL_NETWORK,
|
|
val_metrics['hop_correct'], val_metrics['hop_total'], val_metrics['hop_accuracy'],
|
|
val_metrics['hop_top3_accuracy'],
|
|
val_metrics['complete_paths'], val_metrics['packet_total'],
|
|
val_metrics['complete_path_accuracy'], val_metrics['mean_path_completion'],
|
|
all_metrics['hop_correct'], all_metrics['hop_total'], all_metrics['hop_accuracy'],
|
|
all_metrics['complete_paths'], all_metrics['packet_total'],
|
|
all_metrics['complete_path_accuracy'], json.dumps(params),
|
|
)
|
|
|
|
selection_score = (
|
|
val_metrics['complete_path_accuracy'],
|
|
val_metrics['hop_accuracy'],
|
|
val_metrics['mean_path_completion'],
|
|
all_metrics['complete_path_accuracy'],
|
|
)
|
|
if selection_score > best_score:
|
|
best_score = selection_score
|
|
best_model = model
|
|
best_val_metrics = val_metrics
|
|
best_all_metrics = all_metrics
|
|
best_params = params
|
|
best_rank = rank
|
|
|
|
if best_model is None or best_val_metrics is None or best_all_metrics is None:
|
|
log.warning('All variants failed for global model')
|
|
return
|
|
|
|
log.info(
|
|
'Generation %d winner network=%s variant=%d/%d val_hop=%.3f val_full_path=%.3f val_path_completion=%.3f all_hop=%.3f all_full_path=%.3f (champion val full_path=%.3f all full_path=%.3f)',
|
|
generation, GLOBAL_NETWORK, best_rank, POPULATION_SIZE,
|
|
best_val_metrics['hop_accuracy'], best_val_metrics['complete_path_accuracy'],
|
|
best_val_metrics['mean_path_completion'], best_all_metrics['hop_accuracy'],
|
|
best_all_metrics['complete_path_accuracy'], current_val_best, current_all_best,
|
|
)
|
|
|
|
final_model = train_final_variant(X, y, best_params)
|
|
guard.check()
|
|
if final_model is None:
|
|
log.warning('Final all-gold training failed for generation=%d variant=%d', generation, best_rank)
|
|
final_model = best_model
|
|
|
|
final_all_metrics = evaluate_path_metrics(final_model, X, y, meta_rows)
|
|
if (
|
|
final_all_metrics['complete_path_accuracy'],
|
|
final_all_metrics['hop_accuracy'],
|
|
final_all_metrics['mean_path_completion'],
|
|
) < (
|
|
best_all_metrics['complete_path_accuracy'],
|
|
best_all_metrics['hop_accuracy'],
|
|
best_all_metrics['mean_path_completion'],
|
|
):
|
|
log.warning(
|
|
'Final all-gold model underperformed selected variant for generation=%d variant=%d (final full=%.3f hop=%.3f vs selected full=%.3f hop=%.3f); promoting selected variant model',
|
|
generation, best_rank,
|
|
final_all_metrics['complete_path_accuracy'], final_all_metrics['hop_accuracy'],
|
|
best_all_metrics['complete_path_accuracy'], best_all_metrics['hop_accuracy'],
|
|
)
|
|
final_model = best_model
|
|
final_all_metrics = best_all_metrics
|
|
|
|
log.info(
|
|
'Generation %d final all-gold model network=%s variant=%d/%d gold_replay_hop=%.3f gold_replay_full_path=%.3f gold_replay_path_completion=%.3f heldout_full_path_guardrail=%.3f',
|
|
generation, GLOBAL_NETWORK, best_rank, POPULATION_SIZE,
|
|
final_all_metrics['hop_accuracy'], final_all_metrics['complete_path_accuracy'],
|
|
final_all_metrics['mean_path_completion'], best_val_metrics['complete_path_accuracy'],
|
|
)
|
|
|
|
candidate_guardrail_score = promotion_score(best_val_metrics)
|
|
champion_guardrail_score = (
|
|
promotion_score(champion_val_metrics)
|
|
if champion_val_metrics
|
|
else (0.0, 0.0, 0.0)
|
|
)
|
|
|
|
if champion_val_metrics is None or score_beats(
|
|
candidate_guardrail_score,
|
|
champion_guardrail_score,
|
|
PROMOTION_MIN_DELTA,
|
|
):
|
|
promote_model(
|
|
db, final_model, best_val_metrics, final_all_metrics, int(sum(y)),
|
|
X, y, meta_rows, gold_ids,
|
|
hyperparams=best_params, generation=generation, variant_rank=best_rank,
|
|
guard=guard,
|
|
)
|
|
else:
|
|
log.info(
|
|
'Generation %d network=%s no held-out improvement (candidate val_full=%.3f val_hop=%.3f vs champion val_full=%.3f val_hop=%.3f; candidate all_full=%.3f champion all_full=%.3f), discarding',
|
|
generation, GLOBAL_NETWORK,
|
|
best_val_metrics['complete_path_accuracy'], best_val_metrics['hop_accuracy'],
|
|
champion_guardrail_score[0], champion_guardrail_score[1],
|
|
final_all_metrics['complete_path_accuracy'], current_all_best,
|
|
)
|
|
|
|
|
|
def _request_shutdown(signum, _frame):
|
|
log.info('Received signal %s; requesting graceful shutdown', signum)
|
|
STOP_EVENT.set()
|
|
|
|
|
|
def _install_signal_handlers():
|
|
signal.signal(signal.SIGTERM, _request_shutdown)
|
|
signal.signal(signal.SIGINT, _request_shutdown)
|
|
|
|
|
|
def _rollback_quietly(db):
|
|
try:
|
|
if db is not None and not db.closed:
|
|
db.rollback()
|
|
except Exception:
|
|
pass
|
|
|
|
|
|
def _close_quietly(db):
|
|
try:
|
|
if db is not None and not db.closed:
|
|
db.close()
|
|
except Exception:
|
|
pass
|
|
|
|
|
|
def main():
|
|
log.info(
|
|
'ML path learner starting model_version=%s data_version=%s '
|
|
'batch=%d max_training_rows=%d',
|
|
ML_MODEL_VERSION,
|
|
ML_DATA_VERSION,
|
|
GOLD_BATCH,
|
|
MAX_TRAINING_GOLD_ROWS,
|
|
)
|
|
_install_signal_handlers()
|
|
|
|
while not STOP_EVENT.is_set():
|
|
db = None
|
|
guard = None
|
|
terminal_reason = 'completed'
|
|
training_completed = False
|
|
retry_delay_seconds = GOLD_INTERVAL_SECS
|
|
|
|
try:
|
|
db = get_db()
|
|
guard = claim_leadership(db)
|
|
if guard is None:
|
|
_close_quietly(db)
|
|
db = None
|
|
STOP_EVENT.wait(min(30, max(1, GOLD_INTERVAL_SECS)))
|
|
continue
|
|
|
|
guard.thread.start()
|
|
log.info(
|
|
'Claimed learner lease token=%s training_due=%s',
|
|
guard.token[:8],
|
|
guard.should_train,
|
|
)
|
|
|
|
extracted_rows = 0
|
|
while True:
|
|
guard.check()
|
|
processed_rows = extract_gold_paths(db, guard)
|
|
extracted_rows += processed_rows
|
|
if processed_rows < GOLD_BATCH:
|
|
break
|
|
|
|
if guard.should_train:
|
|
guard.check()
|
|
run_training_cycle(db, guard)
|
|
cleanup_training_artifacts(db, guard)
|
|
training_completed = True
|
|
|
|
terminal_reason = (
|
|
f'completed extraction_rows={extracted_rows} '
|
|
f'training_completed={training_completed}'
|
|
)
|
|
except (KeyboardInterrupt, InterruptedError):
|
|
STOP_EVENT.set()
|
|
terminal_reason = 'graceful shutdown'
|
|
retry_delay_seconds = 0
|
|
_rollback_quietly(db)
|
|
except LeaseLost as exc:
|
|
terminal_reason = f'lease lost: {exc}'
|
|
retry_delay_seconds = 5
|
|
log.warning('%s', terminal_reason)
|
|
_rollback_quietly(db)
|
|
except RunDeadlineExceeded as exc:
|
|
terminal_reason = f'run deadline: {exc}'
|
|
retry_delay_seconds = 30
|
|
log.warning('%s', terminal_reason)
|
|
_rollback_quietly(db)
|
|
except Exception as exc:
|
|
terminal_reason = f'run failed: {type(exc).__name__}: {exc}'
|
|
retry_delay_seconds = 30
|
|
log.error('ML learner run failed: %s', exc, exc_info=True)
|
|
_rollback_quietly(db)
|
|
finally:
|
|
if guard is not None:
|
|
guard.stop_event.set()
|
|
if guard.thread.is_alive():
|
|
guard.thread.join(timeout=5)
|
|
if db is not None and not db.closed:
|
|
try:
|
|
guard.finish(
|
|
db,
|
|
terminal_reason,
|
|
training_completed=training_completed,
|
|
retry_delay_seconds=retry_delay_seconds,
|
|
)
|
|
except Exception as exc:
|
|
log.warning(
|
|
'Could not record learner terminal state: %s',
|
|
exc,
|
|
)
|
|
_rollback_quietly(db)
|
|
_close_quietly(db)
|
|
|
|
if not STOP_EVENT.is_set() and terminal_reason.startswith('run failed'):
|
|
STOP_EVENT.wait(retry_delay_seconds)
|
|
|
|
log.info('ML path learner stopped')
|
|
|
|
|
|
if __name__ == '__main__':
|
|
main()
|