Queimo's picture
Convert Tox21 baseline to TabICLv2
3286d5c verified
Raw History Blame Contribute Delete
3.85 kB
"""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]