Skip to content
Closed
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
137 changes: 118 additions & 19 deletions sdm/models/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,7 @@

import abc
import copy
from collections.abc import Hashable, Iterable, Mapping, Sequence
from collections.abc import Callable, Hashable, Iterable, Mapping, Sequence
from typing import Any, ClassVar, cast

import torch
Expand Down Expand Up @@ -104,6 +104,8 @@ def forward(
recipe: Recipe | None = None,
num_estimators: int | None = None,
estimator_batch_size: int | None = 1,
estimator_cost: Callable[[TableTensor, int], int] | None = None,
estimator_max_cost: int | None = None,
callbacks: Sequence[Callback] | None = None,
generator: torch.Generator | None = None,
**kwargs: Any,
Expand Down Expand Up @@ -132,6 +134,9 @@ def forward(
whose preprocessed tables differ in shape or target class
set, or that come with related tables, run in separate
calls. Device memory grows with the batch size.
estimator_cost: ``(table, num_classes) -> int`` cost of one table.
estimator_max_cost: Run estimators together until this budget
would be exceeded.
callbacks: Callbacks applied in sequence to this model call.
generator: Pseudorandom number generator used for sampling during
pre-processing and model execution.
Expand Down Expand Up @@ -184,6 +189,8 @@ def forward(
contexts=contexts,
queries=queries,
estimator_batch_size=estimator_batch_size,
estimator_cost=estimator_cost,
estimator_max_cost=estimator_max_cost,
callbacks=callbacks,
generator=generator,
**kwargs,
Expand Down Expand Up @@ -212,6 +219,8 @@ def fit(
recipe: Recipe | None = None,
num_estimators: int | None = None,
estimator_batch_size: int | None = 1,
estimator_cost: Callable[[TableTensor, int], int] | None = None,
estimator_max_cost: int | None = None,
callbacks: Sequence[Callback] | None = None,
generator: torch.Generator | None = None,
**kwargs: Any,
Expand Down Expand Up @@ -240,6 +249,9 @@ def fit(
set, or that come with related tables, run in separate
calls. Device memory grows with the batch size. Estimators
fitted together are predicted together.
estimator_cost: ``(table, num_classes) -> int`` cost of one table.
estimator_max_cost: Run estimators together until this budget
would be exceeded.
callbacks: Callbacks applied in sequence to this model call.
generator: Pseudorandom number generator used for sampling during
pre-processing and model execution.
Expand Down Expand Up @@ -273,11 +285,15 @@ def fit(
queries=None,
class_values=class_values,
estimator_batch_size=estimator_batch_size,
max_cost=estimator_max_cost,
cost=estimator_cost,
)
cache = Cache(
recipe_execution=recipe_execution,
kwargs=kwargs,
num_batches=len(batches),
estimator_cost=estimator_cost,
estimator_max_cost=estimator_max_cost,
)
for i, batch in enumerate(batches):
with inference_mode("no_grad"):
Expand Down Expand Up @@ -384,6 +400,13 @@ def predict(
self._cache["recipe_execution"],
)
num_batches = cast(int, self._cache["num_batches"])
estimator_max_cost = (
None if callbacks else self._cache["estimator_max_cost"]
)
estimator_cost = cast(
Callable[[TableTensor, int], int] | None,
self._cache["estimator_cost"],
)
caches = [cast(Cache, self._cache[i]) for i in range(num_batches)]
next_cache = caches[0]

Expand Down Expand Up @@ -418,6 +441,7 @@ def predict(
assert cache is not None

x_schemas = cast(tuple[TableSchema, ...], cache["x_schemas"])
classes = cast(Tensor | None, cache["classes"])
batch_queries = [
self._prepare_query(
query=query,
Expand All @@ -443,20 +467,45 @@ def predict(
with torch.cuda.stream(transfer_stream):
next_cache = next_cache.to(x.device, non_blocking=True)

outs += self._forward_batch(
contexts=None,
queries=batch_queries,
cache=cache,
categorical_mask=cast(Tensor, cache["categorical_mask"]),
class_values=cast(
Sequence[tuple[Any, ...] | None] | None,
cache["class_values"],
),
callbacks=callbacks,
requires_grad=requires_grad,
generator=None,
**cast(dict[str, Any], self._cache["kwargs"]),
)
chunks = [
self._forward_batch(
contexts=None,
queries=chunk,
cache=cache,
categorical_mask=cast(
Tensor, cache["categorical_mask"]
),
class_values=cast(
Sequence[tuple[Any, ...] | None] | None,
cache["class_values"],
),
callbacks=callbacks,
requires_grad=requires_grad,
generator=None,
**cast(dict[str, Any], self._cache["kwargs"]),
)
for chunk in _query_chunks(
queries=batch_queries,
max_cost=cast(int | None, estimator_max_cost),
cost=estimator_cost,
num_classes=(
0 if classes is None else classes.numel()
),
)
]

if len(chunks) == 1:
batch_outputs = chunks[0]
else:
batch_outputs = [
cast(TableTensor, torch.cat(parts, dim=-2))
for parts in zip(*chunks, strict=True)
]

# Drop source chunk references before the next estimator
# batch. A single chunk stays alive through `batch_outputs`.
del chunks
outs.extend(batch_outputs)

if x.is_cuda:
assert compute_stream is not None
Expand Down Expand Up @@ -571,6 +620,8 @@ def _forward_members(
queries: Sequence[MemberQuery],
*,
estimator_batch_size: int | None = 1,
estimator_cost: Callable[[TableTensor, int], int] | None = None,
estimator_max_cost: int | None = None,
callbacks: Sequence[Callback] | None = None,
generator: torch.Generator | None = None,
**kwargs: Any,
Expand All @@ -583,6 +634,9 @@ def _forward_members(
callbacks = () if callbacks is None else callbacks
requires_grad = self.training
requires_grad |= any(callback.requires_grad for callback in callbacks)
if requires_grad and estimator_max_cost is not None:
estimator_batch_size = 1
estimator_max_cost = None

if estimator_batch_size == 1 and len(contexts) > 1:
outs: list[TableTensor] = []
Expand All @@ -592,6 +646,8 @@ def _forward_members(
contexts=(context,),
queries=(query,),
estimator_batch_size=1,
estimator_cost=estimator_cost,
estimator_max_cost=estimator_max_cost,
callbacks=callbacks,
generator=generator,
**kwargs,
Expand Down Expand Up @@ -620,6 +676,8 @@ def _forward_members(
queries=queries,
class_values=class_values,
estimator_batch_size=estimator_batch_size,
max_cost=estimator_max_cost,
cost=estimator_cost,
):
outs += self._forward_batch(
contexts=contexts[batch],
Expand Down Expand Up @@ -849,10 +907,12 @@ def _batch_slices(
queries: Sequence[MemberQuery] | None,
class_values: Sequence[tuple[Any, ...] | None],
estimator_batch_size: int | None,
max_cost: int | None,
cost: Callable[[TableTensor, int], int] | None,
) -> list[slice]:
# Consecutive estimators that can go through one `_forward` together,
# split when shapes/dtypes or the target class set change, or the
# batch is full.
# split when shapes/dtypes or the target class set change, the batch is
# full, or the optional cost budget would be exceeded.
related = any(context.related_tables is not None for context in contexts)
if queries is not None:
related = related or any(
Expand All @@ -862,7 +922,7 @@ def _batch_slices(
return [slice(i, i + 1) for i in range(len(contexts))]

batches: list[slice] = []
start = 0
start = total = 0
key: Hashable = None
# Stack when table shapes/dtypes match and the target class set matches.
for i, context in enumerate(contexts):
Expand All @@ -886,20 +946,59 @@ def _batch_slices(
),
None if classes is None else frozenset(classes),
)
member_cost = 0
if max_cost is not None:
assert cost is not None
num_classes = (
context.y.categorical.categories[0].numel()
if context.y.categorical.size(-1) > 0
else 0
)
member_cost = cost(context.x, num_classes)
if query is not None:
member_cost += cost(query.x, num_classes)
if i > start and (
member_key != key
or (
estimator_batch_size is not None
and i - start == estimator_batch_size
)
or (max_cost is not None and total + member_cost > max_cost)
):
batches.append(slice(start, i))
start = i
start, total = i, 0
key = member_key
total += member_cost
batches.append(slice(start, len(contexts)))
return batches


def _query_chunks(
queries: Sequence[MemberQuery],
max_cost: int | None,
cost: Callable[[TableTensor, int], int] | None,
num_classes: int,
) -> list[list[MemberQuery]]:
# Row chunks of a batch's queries within the cost budget, but never smaller
# than the queries of one estimator on their own.
if max_cost is None or len(queries) == 1:
return [list(queries)]
assert cost is not None
total = sum(cost(query.x, num_classes) for query in queries)
if total <= max_cost:
return [list(queries)]
rows = queries[0].x.size(-2)
budget = max(max_cost, total // len(queries))
splits = [
query.x.split(max(1, budget * rows // total), dim=-2)
for query in queries
]
return [
[MemberQuery(x=x, related_tables=None) for x in xs]
for xs in zip(*splits, strict=True)
]


def _categorical_mask(members: Sequence[MemberContext]) -> Tensor:
x = members[0].x
mask = torch.tensor(
Expand Down
16 changes: 11 additions & 5 deletions sdm/models/ecoc.py
Original file line number Diff line number Diff line change
Expand Up @@ -138,18 +138,24 @@ def forward(
active = index != self.max_classes - 1
return scores.masked_fill(~active, 0).sum(dim=0) / active.sum(dim=0)

def num_tasks(self, num_classes: int) -> int:
"""Return the number of tasks the model runs for ``num_classes``."""
if num_classes <= self.max_classes:
return 1
return max(
# Give every class its own output in at least one task.
math.ceil(num_classes / (self.max_classes - 1)),
4 * math.ceil(math.log(num_classes, self.max_classes)),
)

def _draw_codebook(
self,
num_classes: int,
device: torch.device,
generator: torch.Generator | None,
) -> Tensor:
rest_idx = self.max_classes - 1
num_codes = max(
# Give every class its own output in at least one task.
math.ceil(num_classes / rest_idx),
4 * math.ceil(math.log(num_classes, self.max_classes)),
)
num_codes = self.num_tasks(num_classes)
# Bound the quadratic distance search for large targets.
num_draws = 50 if num_classes <= 200 else 1
codebook = torch.full(
Expand Down
Loading
Loading