"""TabICLv2 classifiers for the twelve Tox21 endpoints.""" from __future__ import annotations import gzip import os import pickle import shutil from typing import Any import numpy as np from tabicl import TabICLClassifier from .utils import TASKS CHECKPOINT_NAME = "model.tabicl.pkl.gz" class Tox21TabICL: """A separate TabICLv2 in-context classifier for each Tox21 endpoint.""" def __init__( self, seed: int = 42, device: str | None = None, model_config: dict[str, Any] | None = None, ): self.tasks = TASKS self.device = device self.model_config = model_config or {} self.model = { task: TabICLClassifier( random_state=seed, device=device, **self.model_config, ) for task in self.tasks } self._shared_model = None def _share_model_weights(self, estimator: TabICLClassifier) -> None: """Keep one immutable TabICLv2 network shared by all task estimators.""" if self._shared_model is None: self._shared_model = estimator.model_ else: estimator.model_ = self._shared_model def load_model(self, path: str) -> None: """Load all fitted task estimators from ``path``.""" for task in self.tasks: model_path = os.path.join(path, task, CHECKPOINT_NAME) with gzip.open(model_path, "rb") as checkpoint_file: estimator = pickle.load(checkpoint_file) if self.device is not None: estimator.device = self.device estimator._resolve_device() estimator.model_.to(estimator.device_) estimator._build_inference_config() self._share_model_weights(estimator) self.model[task] = estimator def save_model(self, path: str) -> None: """Save fitted state while reusing the public TabICLv2 base checkpoint.""" for task in self.tasks: model_path = os.path.join(path, task, CHECKPOINT_NAME) uncompressed_path = model_path.removesuffix(".gz") estimator = self.model[task] # Persist with automatic device selection so checkpoints remain portable. configured_device = estimator.device estimator.device = None try: estimator.save( uncompressed_path, save_model_weights=False, save_training_data=True, save_kv_cache=False, ) with open(uncompressed_path, "rb") as source: with gzip.open(model_path, "wb", compresslevel=1) as destination: shutil.copyfileobj(source, destination) os.remove(uncompressed_path) finally: estimator.device = configured_device def fit(self, task: str, input_features: np.ndarray, labels: np.ndarray) -> None: """Fit the in-context state for one Tox21 endpoint.""" if task not in self.tasks: raise ValueError(f"Unknown task: {task}") if labels.ndim != 1: raise ValueError("Function only accepts one-dimensional labels.") estimator = self.model[task].fit(input_features, labels) self._share_model_weights(estimator) def predict(self, task: str, features: np.ndarray) -> np.ndarray: """Return the probability of the positive class for one endpoint.""" if task not in self.tasks: raise ValueError(f"Unknown task: {task}") if features.ndim != 2: raise ValueError( f"Function expects a two-dimensional array; got {features.shape}." ) return self.model[task].predict_proba(features)[:, 1]