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 +7 -0
- causilo/checkpoints.py +54 -0
- causilo/data/__init__.py +1 -0
- causilo/data/dataset.py +108 -0
- causilo/data/encoding.py +106 -0
- causilo/data/ensemble.py +87 -0
- causilo/data/normalization.py +123 -0
- causilo/engine.py +156 -0
- causilo/estimators.py +168 -0
- causilo/execution/__init__.py +1 -0
- causilo/execution/cached.py +87 -0
- causilo/execution/chunks.py +165 -0
- causilo/execution/direct.py +135 -0
- causilo/execution/memory.py +84 -0
- causilo/execution/precision.py +18 -0
- causilo/execution/runner.py +151 -0
- causilo/model.py +73 -0
- causilo/nn/__init__.py +1 -0
- causilo/nn/cache.py +28 -0
- causilo/nn/column.py +55 -0
- causilo/nn/embeddings.py +50 -0
- causilo/nn/layers/__init__.py +1 -0
- causilo/nn/layers/attention.py +108 -0
- causilo/nn/layers/feedforward.py +17 -0
- causilo/nn/layers/scaling.py +33 -0
- causilo/nn/prediction.py +61 -0
- causilo/nn/row.py +89 -0
- causilo/serialization.py +146 -0
- causilo-1.0.0.dist-info/METADATA +144 -0
- causilo-1.0.0.dist-info/RECORD +33 -0
- causilo-1.0.0.dist-info/WHEEL +5 -0
- causilo-1.0.0.dist-info/licenses/LICENSE +202 -0
- causilo-1.0.0.dist-info/top_level.txt +1 -0
causilo/__init__.py
ADDED
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
|
causilo/data/__init__.py
ADDED
|
@@ -0,0 +1 @@
|
|
|
1
|
+
"""Fitted tabular representations and deterministic ensemble views."""
|
causilo/data/dataset.py
ADDED
|
@@ -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
|
+
)
|
causilo/data/encoding.py
ADDED
|
@@ -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]
|
causilo/data/ensemble.py
ADDED
|
@@ -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
|