causilo 1.0.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
causilo/__init__.py ADDED
@@ -0,0 +1,7 @@
1
+ """Tabular prediction with fixed pretrained context models."""
2
+
3
+ __version__ = "1.0.0"
4
+
5
+ from .estimators import CausiloClassifier, CausiloRegressor
6
+
7
+ __all__ = ["CausiloClassifier", "CausiloRegressor"]
causilo/checkpoints.py ADDED
@@ -0,0 +1,54 @@
1
+ """Load pinned safetensors weights and verify their task/configuration pairing."""
2
+
3
+ import hashlib
4
+ import json
5
+ from pathlib import Path
6
+
7
+ import torch
8
+ from huggingface_hub import snapshot_download
9
+ from safetensors import safe_open
10
+
11
+ from .model import Model, ModelConfig
12
+
13
+ REPOSITORY = "nums-ai/causilo"
14
+ RELEASE_COMMIT = "94f2bd91db0737d4da59f347910662905ecb5a09"
15
+
16
+
17
+ def load_checkpoint(directory: Path) -> Model:
18
+ """Validate the format and config hash, then materialize an evaluation model."""
19
+ encoded = (directory / "config.json").read_bytes()
20
+ record = json.loads(encoded)
21
+ expected = {"architecture": "causilo-v1.0", "format_version": 1, "config_version": 1}
22
+ if any(record.get(key) != value for key, value in expected.items()):
23
+ raise ValueError("Unsupported model configuration")
24
+ config = ModelConfig(**record["model"])
25
+ with safe_open(directory / "model.safetensors", framework="pt") as artifact:
26
+ metadata = artifact.metadata()
27
+ required = {
28
+ "architecture": "causilo-v1.0",
29
+ "format_version": "1",
30
+ "task": config.task,
31
+ "config_sha256": hashlib.sha256(encoded).hexdigest(),
32
+ }
33
+ if metadata != required:
34
+ raise ValueError("Weights do not match their task and configuration")
35
+ tensors = {key: artifact.get_tensor(key) for key in artifact.keys()}
36
+ # Construct shapes without allocating a second full copy of the weights.
37
+ with torch.device("meta"):
38
+ model = Model(config)
39
+ model.load_state_dict(tensors, strict=True, assign=True)
40
+ return model.eval()
41
+
42
+
43
+ def load_pretrained_model(task: str) -> Model:
44
+ """Fetch the task's files at the pinned revision, reusing the Hub cache when available."""
45
+ folder = "classifier" if task == "classification" else "regressor"
46
+ root = snapshot_download(
47
+ REPOSITORY,
48
+ revision=RELEASE_COMMIT,
49
+ allow_patterns=[f"{folder}/config.json", f"{folder}/model.safetensors"],
50
+ )
51
+ model = load_checkpoint(Path(root) / folder)
52
+ if model.config.task != task:
53
+ raise ValueError("The official checkpoint has the wrong prediction task")
54
+ return model
@@ -0,0 +1 @@
1
+ """Fitted tabular representations and deterministic ensemble views."""
@@ -0,0 +1,108 @@
1
+ """Fitted table transforms and ensemble metadata shared by both prediction paths."""
2
+
3
+ from dataclasses import dataclass
4
+
5
+ import numpy as np
6
+ import pandas as pd
7
+ from sklearn.preprocessing import LabelEncoder, StandardScaler
8
+ from sklearn.utils.multiclass import type_of_target
9
+
10
+ from .encoding import FeatureEncoder
11
+ from .ensemble import EnsembleMember, make_ensemble_members
12
+ from .normalization import Normalizer
13
+
14
+
15
+ @dataclass
16
+ class PreparedDataset:
17
+ """Training context in encoded feature space, before member permutations.
18
+
19
+ ``features`` keeps the encoded training table even when transformed tables
20
+ are not retained. ``members`` is in construction order; execution and saved
21
+ K/V caches use ``members_by_normalization()`` order instead. Field names are
22
+ part of the fitted-state serialization format.
23
+ """
24
+
25
+ encoder: FeatureEncoder
26
+ features: np.ndarray
27
+ targets: np.ndarray
28
+ target_encoder: LabelEncoder | StandardScaler
29
+ members: tuple[EnsembleMember, ...]
30
+ normalizers: dict[str, Normalizer]
31
+ normalized_cache: dict[str, np.ndarray]
32
+
33
+ @classmethod
34
+ def prepare(
35
+ cls,
36
+ table,
37
+ targets,
38
+ *,
39
+ task: str,
40
+ n_estimators: int,
41
+ retain_preprocessing: bool,
42
+ max_classes: int,
43
+ random_state: int,
44
+ ) -> "PreparedDataset":
45
+ """Fit transforms on training data only; ``max_classes`` applies to classification."""
46
+ labels = np.asarray(targets)
47
+ if labels.ndim == 2 and labels.shape[1] == 1:
48
+ labels = labels[:, 0]
49
+ if labels.ndim != 1 or len(labels) != len(table):
50
+ raise ValueError("Targets must contain one value per training row")
51
+ if pd.isna(labels).any():
52
+ raise ValueError("Targets cannot contain missing values")
53
+ classes = 0
54
+ if task == "classification":
55
+ if type_of_target(labels) not in {"binary", "multiclass"}:
56
+ raise ValueError("Classification targets must be discrete labels")
57
+ encoder = LabelEncoder().fit(labels)
58
+ classes = len(encoder.classes_)
59
+ if not 1 <= classes <= max_classes:
60
+ raise ValueError(f"Found {classes} classes; this checkpoint supports at most {max_classes}")
61
+ encoded_targets = encoder.transform(labels)
62
+ else:
63
+ # Center in float64 before reducing precision for model execution.
64
+ labels = labels.astype(np.float64)
65
+ if not np.isfinite(labels).all():
66
+ raise ValueError("Regression targets must be finite")
67
+ encoder = StandardScaler().fit(labels[:, None])
68
+ encoded_targets = encoder.transform(labels[:, None])[:, 0].astype(np.float32)
69
+ schema, training = FeatureEncoder.fit(table)
70
+ members = make_ensemble_members(training.shape[1], classes, n_estimators, random_state)
71
+ methods = dict.fromkeys(member.normalization for member in members)
72
+ transforms = {method: Normalizer.fit(training, method) for method in methods}
73
+ retained = (
74
+ {method: transform.transform(training) for method, transform in transforms.items()}
75
+ if retain_preprocessing
76
+ else {}
77
+ )
78
+ return cls(
79
+ encoder=schema,
80
+ features=training,
81
+ targets=encoded_targets,
82
+ target_encoder=encoder,
83
+ members=members,
84
+ normalizers=transforms,
85
+ normalized_cache=retained,
86
+ )
87
+
88
+ def training_table(self, method: str) -> np.ndarray:
89
+ """Return normalized training features, recomputing only if not retained."""
90
+ if method in self.normalized_cache:
91
+ return self.normalized_cache[method]
92
+ return self.normalizers[method].transform(self.features)
93
+
94
+ def targets_for(self, member: EnsembleMember) -> np.ndarray:
95
+ """Map original encoded class IDs to this member's model output IDs."""
96
+ if member.class_order is None:
97
+ return self.targets
98
+ return np.asarray(member.class_order)[self.targets]
99
+
100
+ def members_by_normalization(self) -> tuple[EnsembleMember, ...]:
101
+ """Return the stable execution order, also used when storing K/V caches.
102
+
103
+ Grouping lets members reuse a normalized training table. For eight
104
+ members this returns original indices 0, 4, 1, 5, 2, 6, 3, 7.
105
+ """
106
+ return tuple(
107
+ member for method in self.normalizers for member in self.members if member.normalization == method
108
+ )
@@ -0,0 +1,106 @@
1
+ """Learn a feature schema once and reuse its categories and column selection."""
2
+
3
+ from dataclasses import dataclass
4
+
5
+ import numpy as np
6
+ import pandas as pd
7
+ from sklearn.preprocessing import OrdinalEncoder
8
+
9
+
10
+ @dataclass
11
+ class FeatureEncoder:
12
+ """Persist the input schema and ordinal encoder.
13
+
14
+ ``categories`` and ``continuous`` are input-column indices, not category
15
+ values. Encoded columns are ordered categorical first, then continuous;
16
+ ``retained`` is a varying-feature mask in that encoded order. These field
17
+ names are retained for compatibility with saved fitted states.
18
+ """
19
+
20
+ names: tuple | None
21
+ width: int
22
+ categories: tuple[int, ...]
23
+ continuous: tuple[int, ...]
24
+ encoder: OrdinalEncoder | None
25
+ retained: np.ndarray
26
+
27
+ @classmethod
28
+ def fit(cls, table) -> tuple["FeatureEncoder", np.ndarray]:
29
+ """Infer column kinds and return the schema plus nonconstant encoded features."""
30
+ is_dataframe = isinstance(table, pd.DataFrame)
31
+ array = np.asarray(table)
32
+ if array.ndim != 2 or min(array.shape) == 0:
33
+ raise ValueError("Features must be a nonempty two-dimensional table")
34
+ categorical_indices = []
35
+ for index in range(array.shape[1]):
36
+ dtype = table.dtypes.iloc[index] if is_dataframe else array.dtype
37
+ # Object arrays share one dtype; infer each column without treating
38
+ # numeric-looking strings as continuous measurements.
39
+ if not is_dataframe and pd.api.types.is_object_dtype(dtype):
40
+ inferred = pd.api.types.infer_dtype(array[:, index], skipna=True)
41
+ if inferred in {"integer", "floating", "mixed-integer-float", "decimal"}:
42
+ continue
43
+ if pd.api.types.is_bool_dtype(dtype) or not pd.api.types.is_numeric_dtype(dtype):
44
+ if pd.api.types.is_datetime64_any_dtype(dtype) or pd.api.types.is_timedelta64_dtype(dtype):
45
+ raise TypeError("Date and duration features must be encoded before fitting")
46
+ categorical_indices.append(index)
47
+ continuous_indices = tuple(i for i in range(array.shape[1]) if i not in categorical_indices)
48
+ schema = cls(
49
+ names=tuple(table.columns) if is_dataframe else None,
50
+ width=array.shape[1],
51
+ categories=tuple(categorical_indices),
52
+ continuous=continuous_indices,
53
+ encoder=None,
54
+ retained=np.ones(array.shape[1], dtype=bool),
55
+ )
56
+ if categorical_indices:
57
+ schema.encoder = OrdinalEncoder(
58
+ handle_unknown="use_encoded_value",
59
+ unknown_value=np.nan,
60
+ encoded_missing_value=np.nan,
61
+ dtype=np.float64,
62
+ )
63
+ schema.encoder.fit(schema._categorical(table))
64
+ numeric = schema.encode(table)
65
+ if len(numeric) > 1:
66
+ schema.retained = np.array([np.unique(column).size > 1 for column in numeric.T])
67
+ if not schema.retained.any():
68
+ raise ValueError("At least one varying feature is required")
69
+ return schema, numeric[:, schema.retained]
70
+
71
+ def _categorical(self, table) -> np.ndarray:
72
+ """Unify pandas/NumPy missing markers before ordinal encoding."""
73
+ values = np.asarray(table, dtype=object)[:, self.categories].copy()
74
+ values[pd.isna(values)] = np.nan
75
+ return values
76
+
77
+ def encode(self, table) -> np.ndarray:
78
+ """Validate the input schema and return categorical-first numeric columns.
79
+
80
+ Unknown categories become NaN so the model's missing-value embedding
81
+ handles them. Constant-column removal is applied by ``transform``.
82
+ """
83
+ if isinstance(table, pd.DataFrame):
84
+ if self.names is not None and tuple(table.columns) != self.names:
85
+ raise ValueError("Prediction columns must match fit columns in name and order")
86
+ numerical = table.iloc[:, list(self.continuous)].to_numpy(dtype=np.float64, na_value=np.nan)
87
+ else:
88
+ values = np.asarray(table)
89
+ if values.ndim != 2 or values.shape[1] != self.width:
90
+ raise ValueError(f"Expected {self.width} feature columns")
91
+ numerical = values[:, self.continuous]
92
+ if pd.api.types.is_object_dtype(numerical.dtype):
93
+ numerical = np.where(pd.isna(numerical), np.nan, numerical)
94
+ numerical = np.asarray(numerical, dtype=np.float64)
95
+ if np.asarray(table).shape[1] != self.width:
96
+ raise ValueError(f"Expected {self.width} feature columns")
97
+ parts = [self.encoder.transform(self._categorical(table))] if self.encoder is not None else []
98
+ parts.append(numerical)
99
+ result = np.concatenate(parts, axis=1)
100
+ if np.isinf(result).any():
101
+ raise ValueError("Infinite feature values are not supported")
102
+ return result
103
+
104
+ def transform(self, table) -> np.ndarray:
105
+ """Apply the fitted mappings and varying-feature mask without refitting."""
106
+ return self.encode(table)[:, self.retained]
@@ -0,0 +1,87 @@
1
+ """Deterministic input permutations; all ensemble members share model weights."""
2
+
3
+ import math
4
+ import random
5
+ from dataclasses import dataclass
6
+
7
+ METHOD_ORDER = ("none", "rank2gaussian", "robust", "power")
8
+
9
+
10
+ def _circular_distances(step: int, width: int) -> tuple[int, ...]:
11
+ """Distances to the first two neighbors when walking a ring with this stride."""
12
+ return tuple(
13
+ min((multiple * step) % width, (-multiple * step) % width)
14
+ for multiple in (1, 2)
15
+ )
16
+
17
+
18
+ def make_affine_orders(width: int, count: int, random_state: int) -> list[tuple[int, ...]]:
19
+ """Build feature permutations while discouraging repeated local neighborhoods.
20
+
21
+ Coprime strides visit every feature exactly once. Among at most 32 candidate
22
+ strides, prefer the least-used distances to the first two ring neighbors.
23
+ Keep the first member in identity order. Separate local RNG streams preserve
24
+ the seeded policy without changing Python's global random state.
25
+ """
26
+ if width < 2:
27
+ return [tuple(range(width)) for _ in range(count)]
28
+ generator = random.Random(random_state)
29
+ ring = random.Random(random_state).sample(range(width), width)
30
+ strides = [step for step in range(1, width) if math.gcd(step, width) == 1]
31
+ usage = {}
32
+ orders = []
33
+ for index in range(count):
34
+ candidates = generator.sample(strides, min(32, len(strides))) if index else [1]
35
+ step = min(
36
+ candidates,
37
+ key=lambda stride: sum(usage.get(d, 0) for d in _circular_distances(stride, width)),
38
+ )
39
+ offset = generator.randrange(width) if index else 0
40
+ if index == 0:
41
+ order = tuple(range(width))
42
+ else:
43
+ order = tuple(ring[(offset + step * column) % width] for column in range(width))
44
+ orders.append(order)
45
+ for distance in _circular_distances(step, width):
46
+ usage[distance] = usage.get(distance, 0) + 1
47
+ return orders
48
+
49
+
50
+ @dataclass(frozen=True)
51
+ class EnsembleMember:
52
+ """One normalized/permuted input view of the shared model.
53
+
54
+ ``feature_order`` indexes the encoded, nonconstant feature table.
55
+ ``class_order[original_id]`` is the permuted class ID used as a target and
56
+ the model output column to read when restoring original class order.
57
+ Regression members have no class permutation.
58
+ """
59
+
60
+ normalization: str
61
+ feature_order: tuple[int, ...]
62
+ class_order: tuple[int, ...] | None
63
+
64
+
65
+ def make_ensemble_members(
66
+ width: int, classes: int, count: int, random_state: int
67
+ ) -> tuple[EnsembleMember, ...]:
68
+ """Cycle normalizers and permutations in construction order, fixed by the seed."""
69
+ if count == 1:
70
+ return (EnsembleMember("none", tuple(range(width)), tuple(range(classes)) if classes else None),)
71
+ label_generator = random.Random(random_state)
72
+ class_orders = []
73
+ for index in range(count):
74
+ if not classes:
75
+ class_orders.append(None)
76
+ continue
77
+ if index % classes == 0:
78
+ # Rotate each sampled class ordering so labels occupy every slot
79
+ # before another base ordering is drawn.
80
+ base = label_generator.sample(range(classes), classes)
81
+ offset = index % classes
82
+ class_orders.append(tuple(base[(label + offset) % classes] for label in range(classes)))
83
+ feature_orders = make_affine_orders(width, count, random_state)
84
+ return tuple(
85
+ EnsembleMember(METHOD_ORDER[i % 4], feature_orders[i], class_orders[i])
86
+ for i in range(count)
87
+ )
@@ -0,0 +1,123 @@
1
+ """Training-fitted transforms with explicit missing masks and compressed outliers."""
2
+
3
+ from dataclasses import dataclass
4
+
5
+ import numpy as np
6
+ from scipy.special import ndtri
7
+ from sklearn.preprocessing import PowerTransformer, RobustScaler, StandardScaler
8
+
9
+
10
+ @dataclass
11
+ class Rank2Gaussian:
12
+ """Interpolate training mid-ranks and map them through the Gaussian quantile function."""
13
+
14
+ knots: tuple[tuple[np.ndarray, np.ndarray], ...]
15
+
16
+ @classmethod
17
+ def fit(cls, table):
18
+ """Store unique values and tie-aware probabilities per feature column."""
19
+ knots = []
20
+ for column in table.T:
21
+ values, counts = np.unique(column, return_counts=True)
22
+ probabilities = (np.cumsum(counts) - counts / 2) / len(column)
23
+ knots.append((values, probabilities))
24
+ return cls(tuple(knots))
25
+
26
+ def transform(self, table):
27
+ """Reuse fitted ranks; interpolation clamps unseen extremes to endpoint ranks."""
28
+ return np.column_stack([
29
+ ndtri(np.interp(column, values, probabilities))
30
+ for column, (values, probabilities) in zip(table.T, self.knots)
31
+ ])
32
+
33
+
34
+ @dataclass
35
+ class Normalizer:
36
+ """Normalize encoded features using statistics fitted only on training rows.
37
+
38
+ Missing entries are temporarily mean-filled to fit/apply transforms, then
39
+ restored to NaN for the model's missing-value embedding. Method ``none``
40
+ skips the optional distribution transform, but still standardizes columns
41
+ and compresses outlier tails. Rank-to-Gaussian uses its own final scaling.
42
+ """
43
+
44
+ fill: np.ndarray
45
+ center: np.ndarray
46
+ spread: np.ndarray
47
+ normalizer: object | None
48
+ quantile_scale: StandardScaler | None
49
+ lower: np.ndarray
50
+ upper: np.ndarray
51
+
52
+ @classmethod
53
+ def fit(cls, table: np.ndarray, method: str) -> "Normalizer":
54
+ """Fit filling, scaling, optional distribution transform, and tail bounds."""
55
+ missing = np.isnan(table)
56
+ counts = (~missing).sum(axis=0)
57
+ fill = np.divide(
58
+ np.where(missing, 0, table).sum(axis=0), counts, out=np.zeros(table.shape[1]), where=counts > 0
59
+ )
60
+ filled = np.where(missing, fill, table)
61
+ if method == "rank2gaussian":
62
+ center = np.zeros(filled.shape[1], dtype=filled.dtype)
63
+ spread = np.ones_like(center)
64
+ else:
65
+ center, spread = filled.mean(axis=0), filled.std(axis=0)
66
+ tolerance = np.finfo(filled.dtype).eps * np.maximum(np.abs(center), 1)
67
+ spread = np.where(spread <= tolerance, 1.0, spread)
68
+ scaled = (filled - center) / spread
69
+ normalizer, quantile_scale = None, None
70
+ if method == "power":
71
+ normalizer = PowerTransformer().fit(scaled)
72
+ elif method == "robust":
73
+ normalizer = RobustScaler(unit_variance=True).fit(scaled)
74
+ elif method == "rank2gaussian":
75
+ normalizer = Rank2Gaussian.fit(scaled)
76
+ elif method != "none":
77
+ raise ValueError(f"Unknown internal normalization: {method}")
78
+ normalized = scaled if normalizer is None else normalizer.transform(scaled)
79
+ if method == "rank2gaussian":
80
+ quantile_scale = StandardScaler().fit(normalized)
81
+ normalized = quantile_scale.transform(normalized)
82
+ lower, upper = outlier_bounds(normalized)
83
+ return cls(
84
+ fill=fill,
85
+ center=center,
86
+ spread=spread,
87
+ normalizer=normalizer,
88
+ quantile_scale=quantile_scale,
89
+ lower=lower,
90
+ upper=upper,
91
+ )
92
+
93
+ def transform(self, table: np.ndarray) -> np.ndarray:
94
+ """Return float32 features with the original NaN mask and unchanged fitted state."""
95
+ missing = np.isnan(table)
96
+ values = (np.where(missing, self.fill, table) - self.center) / self.spread
97
+ if self.normalizer is not None:
98
+ values = self.normalizer.transform(values)
99
+ if self.quantile_scale is not None:
100
+ values = self.quantile_scale.transform(values)
101
+ values = compress_tails(values, self.lower, self.upper)
102
+ return np.where(missing, np.nan, values).astype(np.float32)
103
+
104
+
105
+ def outlier_bounds(values):
106
+ """Estimate four-sigma bounds after excluding preliminary extreme values.
107
+
108
+ These thresholds are part of the fixed preprocessing policy. A second pass
109
+ limits how much extreme training values can inflate the final bounds.
110
+ """
111
+ correction = int(len(values) > 1)
112
+ preliminary = np.maximum(values.std(axis=0, ddof=correction), 1e-6)
113
+ mean = values.mean(axis=0)
114
+ central = np.where(np.abs(values - mean) > 4 * preliminary, np.nan, values)
115
+ center = np.nan_to_num(np.nanmean(central, axis=0), nan=0)
116
+ spread = np.maximum(np.nan_to_num(np.nanstd(central, axis=0, ddof=correction), nan=1), 1e-6)
117
+ return center - 4 * spread, center + 4 * spread
118
+
119
+
120
+ def compress_tails(values, lower, upper):
121
+ """Leave central values unchanged and smoothly compress both tails with arcsinh."""
122
+ values = np.where(values < lower, lower - np.arcsinh(lower - values), values)
123
+ return np.where(values > upper, upper + np.arcsinh(values - upper), values)
causilo/engine.py ADDED
@@ -0,0 +1,156 @@
1
+ """Connect fixed pretrained weights to fitted data, caches, and ensemble outputs."""
2
+
3
+ from dataclasses import dataclass
4
+ from numbers import Integral
5
+
6
+ import numpy as np
7
+ import torch
8
+ from sklearn.exceptions import NotFittedError
9
+
10
+ from . import checkpoints
11
+ from .data.dataset import PreparedDataset
12
+ from .data.ensemble import EnsembleMember
13
+ from .execution.cached import cached_predictions
14
+ from .execution.direct import direct_predictions
15
+ from .execution.runner import ModelCache, ModelRunner
16
+
17
+
18
+ def resolve_device(request: str) -> torch.device:
19
+ """Resolve auto placement once; explicit unavailable devices fail instead of falling back."""
20
+ if request == "auto":
21
+ request = "cuda:0" if torch.cuda.is_available() else "cpu"
22
+ device = torch.device(request)
23
+ if device.type not in {"cpu", "cuda"}:
24
+ raise ValueError("Causilo supports CPU or a single CUDA device")
25
+ if device.type == "cuda":
26
+ index = device.index if device.index is not None else 0
27
+ if not torch.cuda.is_available() or index >= torch.cuda.device_count():
28
+ raise ValueError(f"Requested CUDA device {index} is unavailable")
29
+ device = torch.device("cuda", index)
30
+ return device
31
+
32
+
33
+ @dataclass
34
+ class FitState:
35
+ """Successful fit only; caches follow ``dataset.members_by_normalization()``.
36
+
37
+ Keep the stored tuple layout stable for fitted-state restoration. Consumers
38
+ obtain member/cache pairs through ``cached_members`` rather than rebuilding
39
+ the positional association themselves.
40
+ """
41
+
42
+ dataset: PreparedDataset
43
+ caches: tuple[ModelCache, ...] | None
44
+
45
+ def cached_members(self) -> tuple[tuple[EnsembleMember, ModelCache], ...]:
46
+ """Pair every cache with its input permutation, rejecting incomplete state."""
47
+ if self.caches is None:
48
+ raise ValueError("This fit does not contain K/V caches")
49
+ members = self.dataset.members_by_normalization()
50
+ if len(members) != len(self.caches):
51
+ raise ValueError("Cache count does not match the fitted ensemble")
52
+ return tuple(zip(members, self.caches))
53
+
54
+
55
+ class Engine:
56
+ """Own one task/device model and replace its training context on each fit."""
57
+
58
+ def __init__(self, task: str, device: str) -> None:
59
+ self.task = task
60
+ self.device_request = device
61
+ self.device = resolve_device(device)
62
+ self.model = None
63
+ self.state: FitState | None = None
64
+
65
+ def fit(
66
+ self,
67
+ table,
68
+ targets,
69
+ *,
70
+ n_estimators: int,
71
+ use_kv_cache: bool,
72
+ retain_preprocessing: bool,
73
+ random_state: int,
74
+ ) -> None:
75
+ """Prepare context without optimizing weights; a failed fit stays unfitted."""
76
+ self.state = None
77
+ if isinstance(n_estimators, bool) or not isinstance(n_estimators, Integral) or n_estimators < 1:
78
+ raise ValueError("n_estimators must be a positive integer")
79
+ if isinstance(random_state, bool) or not isinstance(random_state, Integral) or random_state < 0:
80
+ raise ValueError("random_state must be a nonnegative integer")
81
+ for name, value in (("use_kv_cache", use_kv_cache), ("retain_preprocessing", retain_preprocessing)):
82
+ if not isinstance(value, (bool, np.bool_)):
83
+ raise ValueError(f"{name} must be a Boolean")
84
+ if self.model is None:
85
+ self.model = checkpoints.load_pretrained_model(self.task).to(self.device)
86
+ prepared = PreparedDataset.prepare(
87
+ table,
88
+ targets,
89
+ task=self.task,
90
+ n_estimators=int(n_estimators),
91
+ retain_preprocessing=retain_preprocessing,
92
+ max_classes=self.model.config.outputs,
93
+ random_state=int(random_state),
94
+ )
95
+ caches = None
96
+ if use_kv_cache:
97
+ runner = ModelRunner(self.model)
98
+ collected = []
99
+ with torch.inference_mode():
100
+ for member in prepared.members_by_normalization():
101
+ features = prepared.training_table(member.normalization)[:, member.feature_order]
102
+ y = prepared.targets_for(member)
103
+ collected.append(runner.build_cache(self._tensor(features), self._tensor(y)))
104
+ caches = tuple(collected)
105
+ # Publish only a complete context; partially built caches stay local.
106
+ self.state = FitState(dataset=prepared, caches=caches)
107
+
108
+ def _tensor(self, value: np.ndarray) -> torch.Tensor:
109
+ """Place one member on the device and prepend its singleton batch axis."""
110
+ return torch.as_tensor(
111
+ np.ascontiguousarray(value), device=self.device, dtype=torch.float32
112
+ ).unsqueeze(0)
113
+
114
+ def require_state(self) -> FitState:
115
+ """Expose fitted state only after all preparation has succeeded."""
116
+ if self.state is None:
117
+ raise NotFittedError("Call fit successfully before prediction")
118
+ return self.state
119
+
120
+ def predict(self, table) -> np.ndarray:
121
+ """Return probabilities (queries, classes) or regression points (queries,)."""
122
+ state = self.require_state()
123
+ fitted = state.dataset
124
+ numeric = fitted.encoder.transform(table)
125
+ outputs = []
126
+
127
+ def reduce_output(result):
128
+ # Preserve the established sorted reduction order for regression.
129
+ # Sorting is algebraically unnecessary, but changing summation order
130
+ # can change floating-point results and saved-state reproducibility.
131
+ return result if self.task == "classification" else result.sort(dim=-1).values.mean(dim=-1)
132
+
133
+ with torch.inference_mode():
134
+ transformed = {
135
+ name: transform.transform(numeric) for name, transform in fitted.normalizers.items()
136
+ }
137
+ if state.caches is None:
138
+ predictions = direct_predictions(self.model, fitted, transformed, self.device, reduce_output)
139
+ else:
140
+ predictions = cached_predictions(
141
+ self.model, fitted, transformed, state.cached_members(), self.device, reduce_output
142
+ )
143
+ for member, result in predictions:
144
+ if self.task == "classification":
145
+ # class_order[original_id] gives the member's output column.
146
+ outputs.append(result[:, member.class_order].numpy())
147
+ else:
148
+ point = result.numpy().astype(np.float64)
149
+ outputs.append(fitted.target_encoder.inverse_transform(point[:, None])[:, 0])
150
+ combined = np.mean(outputs, axis=0)
151
+ if self.task == "classification":
152
+ # Ensemble logits before softmax; averaging probabilities would
153
+ # implement a different prediction policy.
154
+ probabilities = np.exp(combined - combined.max(axis=-1, keepdims=True))
155
+ return probabilities / probabilities.sum(axis=-1, keepdims=True)
156
+ return combined