Spaces:
Running
Running
Download src/model.py from Queimo/tox21_tabicl_classifier: direct link, hf CLI and curl.
- Browser
- Download file 3.85 kB
-
https://huggingface.co/spaces/Queimo/tox21_tabicl_classifier/resolve/main/src/model.py
- Command line
-
hf download hf://spaces/Queimo/tox21_tabicl_classifier/src/model.py
-
curl -L -o model.py https://huggingface.co/spaces/Queimo/tox21_tabicl_classifier/resolve/main/src/model.py
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] | |