File size: 3,853 Bytes
3286d5c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
"""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]