Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
12 changes: 3 additions & 9 deletions benchmark/tabular/model.py
Original file line number Diff line number Diff line change
Expand Up @@ -222,22 +222,16 @@ def _predict_proba(
self.autocast_dtype,
enabled=x_query.is_cuda,
):
outputs = self.model._forward_members(
outputs, dtypes = self.model._forward_members(
contexts=self._contexts,
queries=queries,
estimator_batch_size=self._get_model_params()[
"estimator_batch_size"
],
generator=generator,
)

if self.problem_type == REGRESSION:
outputs = list(
self._recipe_execution.inverse_transform_target(
outputs
)
)
out = self._recipe_execution.transform_output(outputs)
del queries
out = self._recipe_execution.transform_output(outputs, dtypes)

if self.problem_type == REGRESSION:
return out.numerical.float().mean(dim=-1).cpu().numpy()
Expand Down
55 changes: 21 additions & 34 deletions sdm/models/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -195,7 +195,7 @@ def forward(
related_tables=related_query_tables,
)

outs = self._forward_members(
outs, dtypes = self._forward_members(
contexts=contexts,
queries=queries,
estimator_batch_size=estimator_batch_size,
Expand All @@ -204,19 +204,11 @@ def forward(
**kwargs,
)

# Regression: invert target before stacking estimator outputs.
if contexts[0].y.numerical.size(-1) > 0:
with (
torch.amp.autocast(x_query.device.type, enabled=False),
inference_mode("grad" if requires_grad else "inference"),
):
outs = list(recipe_execution.inverse_transform_target(outs))

with (
torch.amp.autocast(x_query.device.type, enabled=False),
inference_mode("grad" if requires_grad else "inference"),
):
return recipe_execution.transform_output(outs)
return recipe_execution.transform_output(outs, dtypes)

def fit(
self,
Expand Down Expand Up @@ -440,6 +432,7 @@ def predict(
compute_stream.wait_stream(transfer_stream)

outs: list[TableTensor] = []
dtypes: list[torch.dtype] = []
start = 0
for i in range(len(caches)):
cache, next_cache = next_cache, None
Expand All @@ -463,6 +456,7 @@ def predict(
strict=True,
)
]
dtypes.extend(query.x.dtype for query in members)
start += len(x_schemas)

if i + 1 < len(caches):
Expand Down Expand Up @@ -527,19 +521,11 @@ def predict(
transfer_stream.synchronize()
raise

# Regression: invert target before stacking estimator outputs.
if cast(Cache, self._cache[0])["classes"] is None:
with (
torch.amp.autocast(x.device.type, enabled=False),
inference_mode("grad" if requires_grad else "inference"),
):
outs = list(recipe_execution.inverse_transform_target(outs))

with (
torch.amp.autocast(x.device.type, enabled=False),
inference_mode("grad" if requires_grad else "inference"),
):
return recipe_execution.transform_output(outs)
return recipe_execution.transform_output(outs, dtypes)

def clear(self) -> None:
r"""Clear cached context state created by :meth:`fit`."""
Expand Down Expand Up @@ -638,13 +624,13 @@ def _forward_members(
callbacks: Sequence[Callback] | None = None,
generator: torch.Generator | None = None,
**kwargs: Any,
) -> list[TableTensor]:
) -> tuple[list[TableTensor], list[torch.dtype]]:
r"""Run recipe-transformed members on the device of their queries.

Context features and targets may live on another device; each batch
is moved when it runs.
Returns one output per member before target inversion and
``recipe.output``.
Returns member outputs before target inversion and ``recipe.output``,
and their callback-prepared query dtypes.
"""
callbacks = () if callbacks is None else callbacks
requires_grad = self.training
Expand All @@ -654,18 +640,19 @@ def _forward_members(

if estimator_batch_size == 1 and len(contexts) > 1:
outs: list[TableTensor] = []
dtypes: list[torch.dtype] = []
for context, query in zip(contexts, queries, strict=True):
outs.extend(
self._forward_members(
contexts=(context,),
queries=(query,),
estimator_batch_size=1,
callbacks=callbacks,
generator=generator,
**kwargs,
)
member_outs, member_dtypes = self._forward_members(
contexts=(context,),
queries=(query,),
estimator_batch_size=1,
callbacks=callbacks,
generator=generator,
**kwargs,
)
return outs
outs.extend(member_outs)
dtypes.extend(member_dtypes)
return outs, dtypes

contexts = [
self._prepare_context(context, callbacks) for context in contexts
Expand Down Expand Up @@ -709,7 +696,7 @@ def _forward_members(
generator=generator,
**kwargs,
)
return outs
return outs, [query.x.dtype for query in queries]

def _prepare_context(
self,
Expand Down Expand Up @@ -788,7 +775,7 @@ def _forward_batch(
for i in range(len(outs)):
for callback in callbacks:
outs[i] = callback.on_model_forward_end(self, outs[i])
return [cast(TableTensor, out.to(query.x.dtype)) for out in outs]
return outs

def _validate_context(
self,
Expand Down
92 changes: 88 additions & 4 deletions sdm/processing/execution.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,14 +4,16 @@
from __future__ import annotations

import copy
from collections.abc import Mapping, Sequence
import math
from collections.abc import Callable, Mapping, Sequence
from typing import NamedTuple, cast

import torch
from torch import Tensor

import sdm.processing as sp
from sdm import EnsembleTable, Recipe, RelatedTables, Stype, TableTensor
from sdm._memory import split_size
from sdm.processing import EnsembleInvertibleMixin, EnsembleProcessor


Expand Down Expand Up @@ -44,6 +46,7 @@ def __init__(self, recipe: Recipe) -> None:
) = None
self._num_estimators: int | None = None
self._y_locations: tuple[tuple[int, int], ...] | None = None
self._numerical_target = False

@property
def num_members(self) -> int:
Expand All @@ -68,6 +71,7 @@ def fit_transform(

self._num_estimators = num_members
self._y_locations = y._locations
self._numerical_target = y[0].numerical.size(-1) > 0

task_dispatchers = tuple(
module
Expand Down Expand Up @@ -166,7 +170,7 @@ def transform(
) -> tuple[MemberQuery, ...]:
"""Transform query data."""
x = _to_ensemble_table(x, self._num_estimators)
x = self.recipe.features.transform_ensemble(x)
x = _transform_rows(self.recipe.features.transform_ensemble, x)
if len(x) != self.num_members:
raise ValueError(
"Expected inputs to map to the same number of ensemble members"
Expand Down Expand Up @@ -254,8 +258,46 @@ def inverse_transform_target(
def transform_output(
self,
outputs: Sequence[TableTensor],
dtypes: Sequence[torch.dtype],
) -> TableTensor:
"""Apply ``recipe.output`` to member outputs."""
"""Apply ``recipe.output`` after restoring each member dtype.

Numerical targets are inverse-transformed before applying the output
recipe.
"""
# Output cells are processed in double precision at most.
size = split_size(
num_items=outputs[0].size(-2),
item_bytes=sum(
math.prod(output.size()[:-2]) * output.size(-1)
for output in outputs
)
* torch.float64.itemsize,
device=outputs[0].device,
)
if size >= outputs[0].size(-2):
return self._transform_output(outputs, dtypes)
parts = tuple(
self._transform_output(chunk, dtypes)
for chunk in zip(
*(output.split(size, dim=-2) for output in outputs),
strict=True,
)
)
return cast(TableTensor, torch.cat(parts, dim=-2))

def _transform_output(
self,
outputs: Sequence[TableTensor],
dtypes: Sequence[torch.dtype],
) -> TableTensor:
outputs = tuple(
cast(TableTensor, output.to(dtype))
for output, dtype in zip(outputs, dtypes, strict=True)
)
if self._numerical_target:
outputs = self.inverse_transform_target(outputs)

if len(outputs) == 1:
out = outputs[0].unsqueeze(0)
else:
Expand All @@ -268,11 +310,53 @@ def transform_output(
"contains the same set of classes."
)

out = torch.stack(list(outputs), dim=0)
out = torch.stack(outputs, dim=0)

return self.recipe.output.transform(cast(TableTensor, out))


def _transform_rows(
transform: Callable[[EnsembleTable], EnsembleTable],
table: EnsembleTable,
) -> EnsembleTable:
# Member cells are transformed in double precision at most.
# Groups have shape [stored members, ..., rows, columns].
num_rows = table._groups[0].size(-2)
row_bytes_by_device: dict[torch.device, int] = {}
# Count logical members because several members may share one stored table.
for group_id, _ in table._locations:
group = table._groups[group_id]
batch_size = math.prod(group.size()[1:-2])
row_bytes = batch_size * group.size(-1) * torch.float64.itemsize
row_bytes_by_device[group.device] = (
row_bytes_by_device.get(group.device, 0) + row_bytes
)
rows_per_pass = min(
split_size(num_rows, row_bytes, device)
for device, row_bytes in row_bytes_by_device.items()
)
if rows_per_pass >= num_rows:
return transform(table)
if any(group.size(-2) != num_rows for group in table._groups[1:]):
raise ValueError(
"Expected all ensemble groups to have the same row count"
)
parts = [
transform(table.replace_groups(groups))
for groups in zip(
*(group.split(rows_per_pass, dim=-2) for group in table._groups),
strict=True,
)
]
# Passes share the fitted state and thus the member layout:
return parts[0].replace_groups(
[
cast(TableTensor, torch.cat(groups, dim=-2))
for groups in zip(*(part._groups for part in parts), strict=True)
]
)


def _align_to_fitted_groups(
ensemble_table: EnsembleTable,
fitted_locations: Sequence[tuple[int, int]],
Expand Down
4 changes: 4 additions & 0 deletions sdm/processing/recipe.py
Original file line number Diff line number Diff line change
Expand Up @@ -29,6 +29,10 @@ class Recipe:
the output remains stacked. Steps before the reducer must support
stacked outputs, while steps after it receive already-reduced outputs.

Fitted ``features`` and ``output`` steps, and inverse ``target`` steps,
transform each row on its own, so models may apply them to large query
sets in passes over rows.

Each pipeline exposes ``fit``/``transform``/``fit_transform`` and, when its
steps are invertible, ``inverse_transform``. Call them directly, e.g.
``recipe.features.transform(table)`` or
Expand Down
49 changes: 48 additions & 1 deletion test/models/test_base.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,7 @@
from sdm.models.callback import Callback
from sdm.processing import InvertibleMixin, Processor
from sdm.processing.execution import RecipeExecution
from sdm.testing import withCUDA


@dataclass
Expand Down Expand Up @@ -320,6 +321,52 @@ def test_model_recipe_generator_does_not_advance_global_rng(
assert torch.equal(torch.get_rng_state(), state)


@withCUDA
@pytest.mark.parametrize("cached", [False, True])
@pytest.mark.parametrize("estimator_batch_size", [1, "auto"])
def test_callback_query_dtype_is_preserved(
device: torch.device,
cached: bool,
estimator_batch_size: int | Literal["auto"],
) -> None:
class DoubleQuery(Callback):
def on_query_preprocessing_end(
self,
model: torch.nn.Module,
x: TableTensor,
related_tables: RelatedTables[TableTensor] | None,
) -> tuple[TableTensor, RelatedTables[TableTensor] | None]:
return (
x.replace_blocks(numerical=x.numerical.double() + 2**-40),
related_tables,
)

model = _RecordingModel().to(device)
context = torch.ones(4, 1, device=device)
target = torch.arange(4, device=device).float().unsqueeze(-1)
query = torch.ones(2, 1, device=device)
if cached:
model.fit(
x=context,
y=target,
num_estimators=2,
estimator_batch_size=estimator_batch_size,
)
output = model.predict(query, callbacks=[DoubleQuery()])
else:
output = model(
x_context=context,
y_context=target,
x_query=query,
num_estimators=2,
estimator_batch_size=estimator_batch_size,
callbacks=[DoubleQuery()],
)

expected = (query.double() + 2**-40).expand(2, -1, -1)
torch.testing.assert_close(output.numerical, expected, rtol=0, atol=0)


def test_callback() -> None:
events: list[str] = []
callbacks = (
Expand Down Expand Up @@ -990,7 +1037,7 @@ def test_forward_members_moves_contexts_to_query_device() -> None:
for query in recipe.transform(x=x_query, related_tables=None)
]

outs = model._forward_members(contexts=contexts, queries=queries)
outs, _ = model._forward_members(contexts=contexts, queries=queries)

assert all(
cast(TableTensor, call.x_context).is_cuda for call in model.calls
Expand Down
Loading
Loading