diff --git a/benchmark/tabular/model.py b/benchmark/tabular/model.py index f363ae972..89b08811a 100644 --- a/benchmark/tabular/model.py +++ b/benchmark/tabular/model.py @@ -130,7 +130,6 @@ def _fit( y_context = y_context[perm].unflatten(0, shape) num_estimators = None self._expand_query = num_estimators is None - self._context_shape = x_context.shape[-2:] recipe = self.model.default_recipe() if params["max_columns"] is not None: @@ -150,7 +149,7 @@ def _fit( y=y_context, recipe=recipe, num_estimators=num_estimators, - estimator_batch_size=self._estimator_batch_size(x_context), + estimator_batch_size=params["estimator_batch_size"], generator=generator, ) return @@ -218,32 +217,19 @@ def _predict_proba( generator = torch.Generator(self._device).set_state( self._rng_state ) - outputs: list[sdm.TableTensor] = [] - size = self._estimator_batch_size(x_query) - size = len(self._contexts) if size is None else size - for start in range(0, len(self._contexts), size): - batch = [ - context._replace( - x=cast( - sdm.TableTensor, context.x.to(self._device) - ), - y=cast( - sdm.TableTensor, context.y.to(self._device) - ), - ) - for context in self._contexts[start : start + size] - ] - with torch.amp.autocast( - self._device.type, - self.autocast_dtype, - enabled=x_query.is_cuda, - ): - outputs += self.model._forward_members( - contexts=batch, - queries=queries[start : start + size], - estimator_batch_size=None, - generator=generator, - ) + with torch.amp.autocast( + self._device.type, + self.autocast_dtype, + enabled=x_query.is_cuda, + ): + outputs = 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( @@ -262,28 +248,6 @@ def _predict_proba( probabilities = out.numerical[..., indices].float().cpu().numpy() return self._convert_proba_to_unified_form(probabilities) - def _estimator_batch_size(self, x: torch.Tensor) -> int | None: - params = self._get_model_params() - estimator_batch_size = params["estimator_batch_size"] - if estimator_batch_size != "auto": - return estimator_batch_size - # Subsampled contexts can produce different cache shapes per estimator. - if not x.is_cuda or self._expand_query: - return 1 - - num_rows, num_cols = self._context_shape - if not params["kv_cache"]: - # Uncached inference processes context and query rows together. - num_rows += x.size(-2) - if num_rows > 3_000 or num_rows * num_cols > 50_000: - return 1 - return self._num_estimators - - num_rows = max(num_rows, x.size(-2)) - if num_rows > 2_000 or num_rows * num_cols >= 50_000: - return 1 - return self._num_estimators - def get_device(self) -> str: return str(next(self.model.parameters()).device) diff --git a/sdm/models/base.py b/sdm/models/base.py index 43228a0b9..dbe5c2c70 100644 --- a/sdm/models/base.py +++ b/sdm/models/base.py @@ -3,8 +3,9 @@ import abc import copy -from collections.abc import Iterable, Mapping, Sequence -from typing import Any, ClassVar, cast +import math +from collections.abc import Callable, Hashable, Iterable, Mapping, Sequence +from typing import Any, ClassVar, Literal, cast import torch from torch import Tensor @@ -60,6 +61,14 @@ class ICLModel(torch.nn.Module, abc.ABC): #: Whether this model supports additional related context. supports_related_tables: ClassVar[bool] + #: Cells per estimator batch with ``estimator_batch_size="auto"``, as + #: counted by :meth:`_estimator_cells`. Larger batches add memory without + #: speeding up inference on GPUs. + _estimator_batch_cells: ClassVar[int] = 2**20 + + #: Cells each table row adds for per-row work such as in-context learning. + _estimator_row_cells: ClassVar[int] = 32 + def __init_subclass__(cls, **kwargs: Any) -> None: super().__init_subclass__(**kwargs) @@ -103,7 +112,7 @@ def forward( *, recipe: Recipe | None = None, num_estimators: int | None = None, - estimator_batch_size: int | None = 1, + estimator_batch_size: int | Literal["auto"] | None = "auto", callbacks: Sequence[Callback] | None = None, generator: torch.Generator | None = None, **kwargs: Any, @@ -126,17 +135,19 @@ def forward( used as the estimator dimension, allowing input data to be customized per estimator (*e.g.*, different in-context examples per estimator). - estimator_batch_size: Maximum number of estimators run through the - model in one call. ``1`` (default) runs estimators one by one; - ``None`` runs all of them together. Estimators in one batch - must share column names, table shapes, category counts and - classes, and related tables require ``1``. Device memory grows - with the batch size. - Model-side randomness drawn per call (*e.g.*, the ECOC codebook - of :class:`~sdm.models.KumoTabular` for more than 10 classes) - is shared within a batch, so batched and sequential predictions - differ numerically there. - callbacks: Callbacks applied in sequence to this model call. + estimator_batch_size: Maximum number of consecutive estimators run + through the model in one call. ``"auto"`` (default) batches + estimators up to a size budget for their preprocessed tables, + and runs them one by one when gradients are required. ``1`` + runs estimators one by one, which minimizes device memory; + ``None`` batches as many as possible. Estimators whose + preprocessed tables differ in shape, category counts or + classes, or that come with related tables, run in separate + calls. Batched and sequential predictions are equal up to + floating-point rounding. + callbacks: Callbacks applied in sequence to this model call. The + preprocessing hooks of all estimators run before their model + forward hooks. generator: Pseudorandom number generator used for sampling during pre-processing and model execution. kwargs: Additional keyword arguments passed to the model. @@ -215,7 +226,7 @@ def fit( *, recipe: Recipe | None = None, num_estimators: int | None = None, - estimator_batch_size: int | None = 1, + estimator_batch_size: int | Literal["auto"] | None = "auto", callbacks: Sequence[Callback] | None = None, generator: torch.Generator | None = None, **kwargs: Any, @@ -237,17 +248,17 @@ def fit( used as the estimator dimension, allowing input data to be customized per estimator (*e.g.*, different in-context examples per estimator). - estimator_batch_size: Maximum number of estimators run through the - model in one call. ``1`` (default) runs estimators one by one; - ``None`` runs all of them together. Estimators in one batch - must share column names, table shapes, category counts and - classes, and related tables require ``1``. Device memory grows - with the batch size. Estimators fitted together are predicted - together. - Model-side randomness drawn per call (*e.g.*, the ECOC codebook - of :class:`~sdm.models.KumoTabular` for more than 10 classes) - is shared within a batch, so batched and sequential predictions - differ numerically there. + estimator_batch_size: Maximum number of consecutive estimators run + through the model in one call. ``"auto"`` (default) batches + estimators up to a size budget for their preprocessed tables. + ``1`` runs estimators one by one, which minimizes device + memory; ``None`` batches as many as possible. Estimators whose + preprocessed tables differ in shape, category counts or + classes, or that come with related tables, run in separate + calls. Estimators fitted together are predicted together; with + ``"auto"`` and without callbacks, :meth:`predict` splits their + query rows into chunks that keep the budget. Batched and + sequential predictions are equal up to floating-point rounding. callbacks: Callbacks applied in sequence to this model call. generator: Pseudorandom number generator used for sampling during pre-processing and model execution. @@ -272,25 +283,37 @@ def fit( generator=generator, ) - if estimator_batch_size is None: - estimator_batch_size = len(contexts) - batches = range(0, len(contexts), estimator_batch_size) + members = [ + self._prepare_context(context, callbacks) for context in contexts + ] + values = _member_class_values(members, estimator_batch_size) + batches = _plan_batches( + contexts=members, + queries=None, + class_values=values, + estimator_batch_size=estimator_batch_size, + max_cells=self._estimator_batch_cells, + cells=self._estimator_cells, + ) cache = Cache( recipe_execution=recipe_execution, kwargs=kwargs, - estimator_batch_size=estimator_batch_size, + num_batches=len(batches), + max_cells=( + self._estimator_batch_cells + if estimator_batch_size == "auto" + else None + ), ) - for i, start in enumerate(batches): - members = [ - self._prepare_context(context, callbacks) - for context in contexts[start : start + estimator_batch_size] - ] + for i, batch in enumerate(batches): with inference_mode("no_grad"): - class_values = _class_values(members) - context = _stack_context(members, class_values) - categorical_mask = _categorical_mask(members) + class_values = _class_values(values[batch]) + context = _stack_context(members[batch]) + categorical_mask = _categorical_mask(members[batch]) batch_cache = Cache( - x_schemas=tuple(member.x.schema for member in members), + x_schemas=tuple( + member.x.schema for member in members[batch] + ), y_schema=context.y.schema, related_tables_schema=context.related_tables.schema if context.related_tables is not None @@ -312,6 +335,7 @@ def fit( cache=batch_cache, generator=generator, categorical_mask=categorical_mask, + num_members=batch.stop - batch.start, **kwargs, ) @@ -387,8 +411,9 @@ def predict( RecipeExecution, self._cache["recipe_execution"], ) - size = cast(int, self._cache["estimator_batch_size"]) - num_batches = len(range(0, recipe_execution.num_members, size)) + num_batches = cast(int, self._cache["num_batches"]) + # Callbacks see every estimator output once, so queries stay whole. + max_cells = None if callbacks else self._cache["max_cells"] caches = [cast(Cache, self._cache[i]) for i in range(num_batches)] next_cache = caches[0] @@ -423,6 +448,7 @@ def predict( assert cache is not None x_schemas = cast(tuple[TableSchema, ...], cache["x_schemas"]) + classes = cast(Tensor | None, cache["classes"]) members = [ self._prepare_query( query=query, @@ -448,19 +474,39 @@ 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=members, - cache=cache, - categorical_mask=cast(Tensor, cache["categorical_mask"]), - class_values=cast( - tuple[tuple[Any, ...], ...] | 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( + tuple[tuple[Any, ...], ...] | 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=members, + max_cells=cast(int | None, max_cells), + cells=self._estimator_cells, + num_classes=( + 0 if classes is None else classes.numel() + ), + ) + ] + outs += ( + chunks[0] + if len(chunks) == 1 + else [ + cast(TableTensor, torch.cat(parts, dim=-2)) + for parts in zip(*chunks, strict=True) + ] ) if x.is_cuda: @@ -526,6 +572,7 @@ def _forward( generator: torch.Generator | None, *, categorical_mask: Tensor, + num_members: int = 1, **kwargs: Any, ) -> TableTensor: # [..., R_query, *] r"""Run the model on preprocessed tables of one estimator batch. @@ -553,6 +600,9 @@ def _forward( numerical feature columns that were categorical before preprocessing. Passed on recording, replaying and uncached calls alike. + num_members: The number of estimators ``E`` in the batch. + Model-side randomness must be drawn per estimator, in order, + so that a batch matches separate calls. kwargs: Additional keyword arguments passed by the caller. Returns: @@ -570,53 +620,76 @@ def default_recipe(cls) -> Recipe: # Helpers ################################################################# + def _estimator_cells(self, x: TableTensor, num_classes: int) -> int: + r"""Count one estimator's table against ``_estimator_batch_cells``. + + Rows of the preprocessed table ``x`` count their columns plus + ``_estimator_row_cells``. ``num_classes`` is ``0`` for regression. + Models that run an estimator as several tasks scale the count up. + """ + rows = math.prod(x.size()[:-1]) + return rows * (x.size(-1) + self._estimator_row_cells) + def _forward_members( self, contexts: Sequence[MemberContext], queries: Sequence[MemberQuery], *, - estimator_batch_size: int | None = 1, + estimator_batch_size: int | Literal["auto"] | None = "auto", callbacks: Sequence[Callback] | None = None, generator: torch.Generator | None = None, **kwargs: Any, ) -> list[TableTensor]: - r"""Run recipe-transformed members that live on the model device. + 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``. """ callbacks = () if callbacks is None else callbacks requires_grad = self.training requires_grad |= any(callback.requires_grad for callback in callbacks) - if estimator_batch_size is None: - estimator_batch_size = len(contexts) + if requires_grad and estimator_batch_size == "auto": + estimator_batch_size = 1 + + members = [ + self._prepare_context(context, callbacks) for context in contexts + ] + queries = [ + self._prepare_query( + query=query, + x_schema=member.x.schema, + related_tables_schema=member.related_tables.schema + if member.related_tables is not None + else None, + callbacks=callbacks, + ) + for member, query in zip(members, queries, strict=True) + ] + values = _member_class_values(members, estimator_batch_size) outs: list[TableTensor] = [] - for start in range(0, len(contexts), estimator_batch_size): - members = [ - self._prepare_context(context, callbacks) - for context in contexts[start : start + estimator_batch_size] - ] - query_members = [ - self._prepare_query( - query=query, - x_schema=member.x.schema, - related_tables_schema=member.related_tables.schema - if member.related_tables is not None - else None, - callbacks=callbacks, - ) - for member, query in zip( - members, - queries[start : start + estimator_batch_size], - strict=True, - ) - ] + for batch in _plan_batches( + contexts=members, + queries=queries, + class_values=values, + estimator_batch_size=estimator_batch_size, + max_cells=self._estimator_batch_cells, + cells=self._estimator_cells, + ): + device = queries[batch.start].x.device outs += self._forward_batch( - contexts=members, - queries=query_members, + contexts=[ + member._replace( + x=cast(TableTensor, member.x.to(device)), + y=cast(TableTensor, member.y.to(device)), + ) + for member in members[batch] + ], + queries=queries[batch], cache=None, categorical_mask=None, - class_values=_class_values(members), + class_values=_class_values(values[batch]), callbacks=callbacks, requires_grad=requires_grad, generator=generator, @@ -679,11 +752,7 @@ def _forward_batch( # Stacking inside the autograd region keeps callback-captured leaves # attached to the graph. with inference_mode("grad" if requires_grad else "inference"): - context = ( - None - if contexts is None - else _stack_context(contexts, class_values) - ) + context = None if contexts is None else _stack_context(contexts) if categorical_mask is None: assert contexts is not None categorical_mask = _categorical_mask(contexts) @@ -699,6 +768,7 @@ def _forward_batch( cache=cache, generator=generator, categorical_mask=categorical_mask, + num_members=len(queries), **kwargs, ) outs = _unstack(out, class_values, len(queries)) @@ -819,44 +889,12 @@ def _stack(tables: Sequence[TableTensor]) -> TableTensor: return cast(TableTensor, torch.stack(renamed)) -_INCOMPATIBLE_ESTIMATORS = ( - "Estimators in one batch must share column names, category counts and " - "classes; use 'estimator_batch_size=1'" -) - - -def _check_compatible(tables: Sequence[TableTensor]) -> None: - ref = tables[0] - names = {stype: frozenset(names) for stype, names in ref.columns.items()} - counts = tuple(c.numel() for c in ref.categorical.categories) - for table in tables[1:]: - if { - stype: frozenset(names) for stype, names in table.columns.items() - } != names or ( - tuple(c.numel() for c in table.categorical.categories) != counts - ): - raise ValueError(_INCOMPATIBLE_ESTIMATORS) - - -def _stack_context( - members: Sequence[MemberContext], - class_values: tuple[tuple[Any, ...], ...] | None, -) -> MemberContext: +def _stack_context(members: Sequence[MemberContext]) -> MemberContext: if len(members) == 1: return members[0] - if members[0].related_tables is not None: - raise ValueError("Related tables require 'estimator_batch_size=1'") - xs = [member.x for member in members] - ys = [member.y for member in members] - _check_compatible(xs) - _check_compatible(ys) - if class_values is not None and any( - set(values) != set(class_values[0]) for values in class_values[1:] - ): - raise ValueError(_INCOMPATIBLE_ESTIMATORS) return MemberContext( - x=_stack(xs), - y=_stack(ys), + x=_stack([member.x for member in members]), + y=_stack([member.y for member in members]), related_tables=None, input_stypes=members[0].input_stypes, ) @@ -865,9 +903,106 @@ def _stack_context( def _stack_query(members: Sequence[MemberQuery]) -> MemberQuery: if len(members) == 1: return members[0] - xs = [member.x for member in members] - _check_compatible(xs) - return MemberQuery(x=_stack(xs), related_tables=None) + return MemberQuery( + x=_stack([member.x for member in members]), related_tables=None + ) + + +def _plan_batches( + contexts: Sequence[MemberContext], + queries: Sequence[MemberQuery] | None, + class_values: Sequence[tuple[Any, ...] | None], + estimator_batch_size: int | Literal["auto"] | None, + max_cells: int, + cells: Callable[[TableTensor, int], int], +) -> list[slice]: + # Runs of consecutive members whose tables stack, cut at the batch size or, + # for "auto", before the cell budget is exceeded. Keeping members in order + # keeps model-side randomness in the order of separate calls. + batches: list[slice] = [] + start = total = 0 + key: Hashable = None + for i, context in enumerate(contexts): + query = None if queries is None else queries[i] + member_key = _stack_key(context, query, class_values[i]) + num_classes = _num_classes(context) + member_cells = cells(context.x, num_classes) + if query is not None: + member_cells += cells(query.x, num_classes) + if i > start and ( + member_key is None + or member_key != key + or i - start == estimator_batch_size + or ( + estimator_batch_size == "auto" + and total + member_cells > max_cells + ) + ): + batches.append(slice(start, i)) + start, total = i, 0 + key = member_key + total += member_cells + batches.append(slice(start, len(contexts))) + return batches + + +def _stack_key( + context: MemberContext, + query: MemberQuery | None, + class_values: tuple[Any, ...] | None, +) -> Hashable: + # Members stack when their tables agree in block shapes, dtypes and + # category counts and their targets in classes; ``None`` never stacks. + if context.related_tables is not None or ( + query is not None and query.related_tables is not None + ): + return None + tables = ( + [context.x, context.y] + if query is None + else [context.x, context.y, query.x] + ) + return ( + tuple( + ( + tuple( + (stype, block.size(), block.dtype) + for stype, block in table.items() + ), + tuple(c.numel() for c in table.categorical.categories), + ) + for table in tables + ), + None if class_values is None else frozenset(class_values), + ) + + +def _num_classes(member: MemberContext) -> int: + y = member.y.categorical + return y.categories[0].numel() if y.size(-1) > 0 else 0 + + +def _query_chunks( + queries: Sequence[MemberQuery], + max_cells: int | None, + cells: Callable[[TableTensor, int], int], + num_classes: int, +) -> list[list[MemberQuery]]: + # Row chunks of a batch's queries within the cell budget, but never smaller + # than the queries of one estimator on their own. + total = sum(cells(query.x, num_classes) for query in queries) + if max_cells is None or len(queries) == 1 or total <= max_cells: + return [list(queries)] + rows = queries[0].x.size(-2) + budget = max(max_cells, 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: @@ -889,15 +1024,25 @@ def _categorical_mask(members: Sequence[MemberContext]) -> Tensor: return mask.view(len(members), *(1,) * (x.dim() - 2), -1) # [E, 1, ..., C] -def _class_values( +def _member_class_values( members: Sequence[MemberContext], -) -> tuple[tuple[Any, ...], ...] | None: - if len(members) == 1 or members[0].y.categorical.size(-1) == 0: - return None - return tuple( + estimator_batch_size: int | Literal["auto"] | None, +) -> list[tuple[Any, ...] | None]: + # Classes of every member, read only when members may share a batch. + if estimator_batch_size == 1 or members[0].y.categorical.size(-1) == 0: + return [None] * len(members) + return [ tuple(member.y.categorical.categories[0].tolist()) for member in members - ) + ] + + +def _class_values( + values: Sequence[tuple[Any, ...] | None], +) -> tuple[tuple[Any, ...], ...] | None: + if len(values) == 1 or values[0] is None: + return None + return cast(tuple[tuple[Any, ...], ...], tuple(values)) def _unstack( diff --git a/sdm/models/ecoc.py b/sdm/models/ecoc.py index 5fe14e72c..daf6df26e 100644 --- a/sdm/models/ecoc.py +++ b/sdm/models/ecoc.py @@ -49,6 +49,7 @@ def forward( y: Tensor, *, num_classes: int, + num_members: int = 1, cache: Cache | None = None, generator: torch.Generator | None = None, **kwargs: Any, @@ -65,6 +66,9 @@ def forward( ``[..., R_context]``. num_classes: Total number of target classes ``K``, including those absent from the context. + num_members: Number of ensemble members stacked along the first + dimension of ``x`` and ``y``. Each member draws its own + codebook, in member order, as separate calls would. cache: Cache for model state and the codebook. On replay, pass query-only ``x``, an empty context axis in ``y``, and the same ``num_classes``. @@ -84,29 +88,49 @@ def forward( if cache is not None and cache.is_replaying: codebook = cast(Tensor, cache["ecoc_codebook"]) - if num_classes != codebook.size(1): + if num_classes != codebook.size(-1): raise ValueError( "'num_classes' must match the cached ECOC codebook " - f"(expected {codebook.size(1)}, got {num_classes})" + f"(expected {codebook.size(-1)}, got {num_classes})" ) kwargs["cache"] = cast(Cache, cache["ecoc_model"]) - else: + elif num_members == 1: # [T, K] codebook = self._draw_codebook(num_classes, x.device, generator) - if cache is not None: - cache["ecoc_codebook"] = codebook - kwargs["cache"] = cache["ecoc_model"] = Cache() - - T = codebook.size(0) - logits = model( - x=x.expand(T, *x.shape), # [T, ..., R, C] - y=codebook.index_select( - dim=1, - index=y.reshape(-1), - ).view(T, *y.shape), # [T, ..., R_context] - **kwargs, - ) - index = codebook.view(T, *(1,) * (logits.dim() - 2), num_classes) + else: + # [E, T, K] + codebook = torch.stack( + [ + self._draw_codebook(num_classes, x.device, generator) + for _ in range(num_members) + ] + ) + if cache is not None and not cache.is_replaying: + cache["ecoc_codebook"] = codebook + kwargs["cache"] = cache["ecoc_model"] = Cache() + + T = codebook.size(-2) + if codebook.dim() == 2: + y = codebook.index_select(dim=1, index=y.reshape(-1)).view( + T, *y.shape + ) # [T, ..., R_context] + else: + E = codebook.size(0) + y = ( + codebook.gather( + dim=-1, + index=y.long().reshape(E, 1, -1).expand(E, T, -1), + ) # [E, T, N] + .view(E, T, *y.shape[1:]) + .movedim(1, 0) + ) # [T, E, ..., R_context] + logits = model(x=x.expand(T, *x.shape), y=y, **kwargs) + if codebook.dim() == 2: + index = codebook.view(T, *(1,) * (logits.dim() - 2), num_classes) + else: + index = codebook.movedim(1, 0).reshape( + T, E, *(1,) * (logits.dim() - 3), num_classes + ) scores = logits.log_softmax(dim=-1).gather( dim=-1, index=index.expand(*logits.shape[:-1], num_classes), @@ -114,6 +138,16 @@ 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, @@ -121,11 +155,7 @@ def _draw_codebook( 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( diff --git a/sdm/models/kumo/tabular/model.py b/sdm/models/kumo/tabular/model.py index 814bbc9ab..204fdd95e 100644 --- a/sdm/models/kumo/tabular/model.py +++ b/sdm/models/kumo/tabular/model.py @@ -176,6 +176,7 @@ def _forward( generator: torch.Generator | None, *, categorical_mask: Tensor, + num_members: int = 1, **kwargs: Any, ) -> TableTensor: # [..., R_query, num_classes or 999] @@ -223,6 +224,7 @@ def _forward( x=x, y=y, num_classes=len(classes), + num_members=num_members, cache=cache, generator=generator, categorical_mask=categorical_mask, @@ -232,6 +234,12 @@ def _forward( numerical=out, ) + def _estimator_cells(self, x: TableTensor, num_classes: int) -> int: + cells = super()._estimator_cells(x, num_classes) + if num_classes == 0: + return cells + return cells * self.ecoc.num_tasks(num_classes) + class _KumoTabular(torch.nn.Module): def __init__( diff --git a/sdm/nn/attention.py b/sdm/nn/attention.py index b2020d9ff..b0d71c09c 100644 --- a/sdm/nn/attention.py +++ b/sdm/nn/attention.py @@ -649,7 +649,7 @@ def forward( for size in reversed(batch_shape): batch_indices.append(flat_index % size) flat_index = flat_index // size - out[tuple(reversed(batch_indices))] = chunk + out[tuple(reversed(batch_indices))] = chunk.to(out.dtype) elif flat_out is None: flat_out = chunk.new_empty((batch_size, *query.size()[-2:])) flat_out[start:end] = chunk diff --git a/test/models/kumo/tabular/test_model.py b/test/models/kumo/tabular/test_model.py index 48b2cf008..71e1782bc 100644 --- a/test/models/kumo/tabular/test_model.py +++ b/test/models/kumo/tabular/test_model.py @@ -131,11 +131,11 @@ def test_categorical_features_are_marked(cls_model: KumoTabular) -> None: @pytest.mark.parametrize("task", ["classification", "regression"]) -@pytest.mark.parametrize("estimator_batch_size", [2, None]) +@pytest.mark.parametrize("estimator_batch_size", [2, None, "auto"]) def test_estimator_batching( task: Literal["classification", "regression"], size: Literal["small", "medium", "large"], - estimator_batch_size: int | None, + estimator_batch_size: int | Literal["auto"] | None, ) -> None: model = _build(task, size) x_context, x_query = _features() @@ -146,6 +146,7 @@ def test_estimator_batching( y_context=target, x_query=x_query, num_estimators=5, + estimator_batch_size=1, generator=torch.Generator().manual_seed(0), ) @@ -258,3 +259,120 @@ def test_missing_values_pass_through_fit_predict( assert cached.shape == direct.shape assert (direct.numerical.diff(dim=-1) >= 0).all() assert (cached.numerical.diff(dim=-1) >= 0).all() + + +def _assert_batching_matches_sequential( + model: KumoTabular, + x_context: TableTensor, + y_context: TableTensor, + x_query: TableTensor, + estimator_batch_size: int | Literal["auto"] | None, +) -> None: + def forward(size: int | Literal["auto"] | None) -> TableTensor: + return model( + x_context=x_context, + y_context=y_context, + x_query=x_query, + num_estimators=4, + estimator_batch_size=size, + generator=torch.Generator().manual_seed(0), + ) + + expected = forward(1) + actual = forward(estimator_batch_size) + assert actual.columns == expected.columns + torch.testing.assert_close(actual.numerical, expected.numerical) + + model.fit( + x=x_context, + y=y_context, + num_estimators=4, + estimator_batch_size=estimator_batch_size, + generator=torch.Generator().manual_seed(0), + ) + actual = model.predict(x_query) + assert actual.columns == expected.columns + torch.testing.assert_close(actual.numerical, expected.numerical) + + +@pytest.mark.parametrize("task", ["classification", "regression"]) +@pytest.mark.parametrize("estimator_batch_size", [None, "auto"]) +def test_estimator_batching_wide_table( + task: Literal["classification", "regression"], + estimator_batch_size: Literal["auto"] | None, +) -> None: + # The default recipe keeps a different set of 500 columns per estimator. + model = _build(task, "small") + x_context, x_query = TableTensor.from_tensor(torch.randn(12, 520)).split( + 8, dim=0 + ) + target = ( + TableTensor( + columns={Stype.categorical: ("target",)}, + categorical=CategoricalTensor( + code=torch.tensor([[0], [1], [2], [0], [1], [2], [0], [1]]), + categories=(torch.arange(3),), + ), + ) + if task == "classification" + else TableTensor.from_tensor(torch.randn(8, 1)) + ) + _assert_batching_matches_sequential( + model=model, + x_context=x_context, + y_context=target, + x_query=x_query, + estimator_batch_size=estimator_batch_size, + ) + + +@pytest.mark.parametrize("estimator_batch_size", [2, None, "auto"]) +def test_estimator_batching_many_classes( + estimator_batch_size: int | Literal["auto"] | None, +) -> None: + # More than 10 classes run through ECOC codebooks drawn per estimator. + model = _build("classification", "small") + x_context, x_query = TableTensor.from_tensor(torch.randn(28, 4)).split( + 24, dim=0 + ) + target = TableTensor( + columns={Stype.categorical: ("target",)}, + categorical=CategoricalTensor( + code=torch.arange(24).remainder(12).unsqueeze(-1), + categories=(torch.arange(12),), + ), + ) + _assert_batching_matches_sequential( + model=model, + x_context=x_context, + y_context=target, + x_query=x_query, + estimator_batch_size=estimator_batch_size, + ) + + +def test_estimator_batching_many_classes_query_chunks( + monkeypatch: pytest.MonkeyPatch, +) -> None: + # 4 estimators of 24 context rows x 4 columns x 8 ECOC tasks fit one batch, + # so predict replays their codebooks over chunks of the 40 query rows. + monkeypatch.setattr(KumoTabular, "_estimator_batch_cells", 4 * 24 * 4 * 8) + monkeypatch.setattr(KumoTabular, "_estimator_row_cells", 0) + model = _build("classification", "small") + x_context, x_query = TableTensor.from_tensor(torch.randn(64, 4)).split( + [24, 40], dim=0 + ) + target = TableTensor( + columns={Stype.categorical: ("target",)}, + categorical=CategoricalTensor( + code=torch.arange(24).remainder(12).unsqueeze(-1), + categories=(torch.arange(12),), + ), + ) + _assert_batching_matches_sequential( + model=model, + x_context=x_context, + y_context=target, + x_query=x_query, + estimator_batch_size="auto", + ) diff --git a/test/models/tabfm/test_model.py b/test/models/tabfm/test_model.py index 850a686ea..919b8772f 100644 --- a/test/models/tabfm/test_model.py +++ b/test/models/tabfm/test_model.py @@ -2,6 +2,7 @@ # SPDX-License-Identifier: Apache-2.0 import functools +from typing import Literal import pytest import torch @@ -14,11 +15,11 @@ @withCUDA @pytest.mark.parametrize("dtype", [torch.int64, torch.float32]) -@pytest.mark.parametrize("estimator_batch_size", [1, 2, None]) +@pytest.mark.parametrize("estimator_batch_size", [1, 2, None, "auto"]) def test_forward( device: torch.device, dtype: torch.dtype, - estimator_batch_size: int | None, + estimator_batch_size: int | Literal["auto"] | None, monkeypatch: pytest.MonkeyPatch, ) -> None: monkeypatch.setattr( @@ -66,6 +67,7 @@ def test_forward( y_context=y_context, x_query=x_query, num_estimators=9, + estimator_batch_size=1, generator=generator, ) assert out.dtype == x_context.dtype @@ -99,3 +101,55 @@ def test_forward( assert model._cache.size() > 0 assert model.predict(x_query).allclose(out, atol=1e-4, rtol=1e-4) model.clear() + + +@pytest.mark.parametrize("estimator_batch_size", [None, "auto"]) +def test_estimator_batching_wide_table( + estimator_batch_size: Literal["auto"] | None, + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr( + tabfm_module, + "_TabFM", + functools.partial( + tabfm_module._TabFM, + channels=64, + num_inducing_points=128, + num_readout_tokens=4, + num_icl_layers=4, + ), + ) + model = TabFM(task="classification", pretrained=False) + for parameter in model.parameters(): + if not parameter.any(): + torch.nn.init.normal_(parameter, std=0.02) + # The default recipe keeps a different set of 500 columns per estimator. + x_context, x_query = TableTensor.from_tensor(torch.randn(12, 520)).split( + 8, dim=0 + ) + y_context = torch.tensor([0, 1, 0, 1, 0, 1, 0, 1]).unsqueeze(-1) + + def forward(size: int | Literal["auto"] | None) -> TableTensor: + return model( + x_context=x_context, + y_context=y_context, + x_query=x_query, + num_estimators=4, + estimator_batch_size=size, + generator=torch.Generator().manual_seed(1), + ) + + expected = forward(1) + torch.testing.assert_close( + forward(estimator_batch_size).numerical, expected.numerical + ) + model.fit( + x=x_context, + y=y_context, + num_estimators=4, + estimator_batch_size=estimator_batch_size, + generator=torch.Generator().manual_seed(1), + ) + torch.testing.assert_close( + model.predict(x_query).numerical, expected.numerical + ) diff --git a/test/models/tabiclv2/test_model.py b/test/models/tabiclv2/test_model.py index 43f98ddfe..77137798e 100644 --- a/test/models/tabiclv2/test_model.py +++ b/test/models/tabiclv2/test_model.py @@ -154,7 +154,7 @@ def test_num_estimators( ) assert out.size() == (*batch_shape, R_query, 999) - model.fit(x_context, y_context, num_estimators=3) + model.fit(x_context, y_context, num_estimators=3, estimator_batch_size=1) assert model._cache is not None assert 0 in model._cache assert 1 in model._cache diff --git a/test/models/test_base.py b/test/models/test_base.py index b5ad6205e..6c5171507 100644 --- a/test/models/test_base.py +++ b/test/models/test_base.py @@ -2,7 +2,7 @@ # SPDX-License-Identifier: Apache-2.0 from dataclasses import dataclass -from typing import Any, ClassVar, cast +from typing import Any, ClassVar, Literal, cast import pytest import torch @@ -20,6 +20,7 @@ from sdm.models import ICLModel from sdm.models.callback import Callback from sdm.processing import InvertibleMixin, Processor +from sdm.processing.execution import RecipeExecution @dataclass @@ -79,6 +80,7 @@ class _ClassFrequencyModel(ICLModel): def __init__(self) -> None: super().__init__(task=None) + self.num_calls = 0 self.eval() def _forward( @@ -92,6 +94,7 @@ def _forward( generator: torch.Generator | None, **kwargs: Any, ) -> TableTensor: + self.num_calls += 1 if cache is None or cache.is_recording: assert y_context is not None classes = y_context.categorical.categories[0] @@ -604,8 +607,10 @@ def test_ensemble_output_reduce() -> None: assert out.size() == (2, 3) -@pytest.mark.parametrize("estimator_batch_size", [1, 2, None]) -def test_estimator_callbacks(estimator_batch_size: int | None) -> None: +@pytest.mark.parametrize("estimator_batch_size", [1, 2, None, "auto"]) +def test_estimator_callbacks( + estimator_batch_size: int | Literal["auto"] | None, +) -> None: model = _RecordingModel() x = torch.arange(30.0).view(5, 3, 2) y = torch.zeros(5, 3, 1) @@ -643,10 +648,15 @@ def test_estimator_callbacks(estimator_batch_size: int | None) -> None: @pytest.mark.parametrize( ("estimator_batch_size", "num_calls", "query_size"), - [(1, 4, (2, 2)), (2, 2, (2, 2, 2)), (None, 1, (4, 2, 2))], + [ + (1, 4, (2, 2)), + (2, 2, (2, 2, 2)), + (None, 1, (4, 2, 2)), + ("auto", 1, (4, 2, 2)), + ], ) def test_estimator_batching_groups_consecutive_members( - estimator_batch_size: int | None, + estimator_batch_size: int | Literal["auto"] | None, num_calls: int, query_size: tuple[int, ...], ) -> None: @@ -722,7 +732,7 @@ def _check(out: TableTensor) -> None: _check(model.predict(x)) -def test_estimator_batching_rejects_related_tables() -> None: +def test_estimator_batching_runs_related_tables_one_by_one() -> None: model = _RecordingModel() x_context = _table([0.0, 2.0], [1, 2], value_column="feature") x_query = _table([3.0], [3], value_column="feature") @@ -730,27 +740,34 @@ def test_estimator_batching_rejects_related_tables() -> None: related_context = _related_tables(query=False) related_query = _related_tables(query=True) - with pytest.raises(ValueError, match="Related tables require"): - model( + def forward(size: int | None) -> TableTensor: + return model( x_context=x_context, y_context=y_context, x_query=x_query, related_context_tables=related_context, related_query_tables=related_query, num_estimators=2, - estimator_batch_size=None, - ) - with pytest.raises(ValueError, match="Related tables require"): - model.fit( - x=x_context, - y=y_context, - related_tables=related_context, - num_estimators=2, - estimator_batch_size=None, + estimator_batch_size=size, ) + expected = forward(1) + model.calls.clear() + torch.testing.assert_close(forward(None).numerical, expected.numerical) + assert len(model.calls) == 2 -def test_estimator_batching_incompatible_shapes() -> None: + model.fit( + x=x_context, + y=y_context, + related_tables=related_context, + num_estimators=2, + estimator_batch_size=None, + ) + actual = model.predict(x_query, related_query) + torch.testing.assert_close(actual.numerical, expected.numerical) + + +def test_estimator_batching_splits_different_shapes() -> None: model = _RecordingModel() x = EnsembleTable.from_tables( tables=[ @@ -766,16 +783,17 @@ def test_estimator_batching_incompatible_shapes() -> None: ], member_table_ids=(0, 1), ) - with pytest.raises(RuntimeError, match="stack expects"): - model.fit(x, y, estimator_batch_size=None) query = TableTensor(numerical=torch.randn(2, 2, 2)) - with pytest.raises(RuntimeError, match="stack expects"): - model(x, y, query, estimator_batch_size=None) - model.fit(x, y) + + out = model(x, y, query, estimator_batch_size=None) + torch.testing.assert_close(out.numerical, query.numerical) + assert len(model.calls) == 2 + + model.fit(x, y, estimator_batch_size=None) torch.testing.assert_close(model.predict(query).numerical, query.numerical) -def test_estimator_batching_incompatible_categories() -> None: +def test_estimator_batching_splits_different_category_counts() -> None: model = _RecordingModel() y = EnsembleTable.from_tables( tables=[ @@ -789,51 +807,76 @@ def test_estimator_batching_incompatible_categories() -> None: ], member_table_ids=(0, 1), ) - with pytest.raises(ValueError, match="Estimators in one batch"): - model.fit( - x=torch.ones(3, 2), - y=y, - num_estimators=2, - estimator_batch_size=None, - ) - with pytest.raises(ValueError, match="Estimators in one batch"): - model( - x_context=torch.ones(3, 2), - y_context=y, - x_query=torch.ones(1, 2), - num_estimators=2, - estimator_batch_size=None, - ) + query = torch.randn(1, 2) + + out = model( + x_context=torch.ones(3, 2), + y_context=y, + x_query=query, + num_estimators=2, + estimator_batch_size=None, + ) + torch.testing.assert_close(out.numerical, query.expand(2, 1, 2)) + assert len(model.calls) == 2 + + model.fit( + x=torch.ones(3, 2), + y=y, + num_estimators=2, + estimator_batch_size=None, + ) + torch.testing.assert_close( + model.predict(query).numerical, + query.expand(2, 1, 2), + ) -@pytest.mark.parametrize("columns", [("a", "c"), ("a",)]) -def test_estimator_batching_incompatible_columns( +@pytest.mark.parametrize( + ("columns", "num_calls"), + [(("a", "c"), 1), (("a",), 2)], +) +def test_estimator_batching_stacks_tables_of_equal_shape( columns: tuple[str, ...], + num_calls: int, ) -> None: - model = _RecordingModel() + # Column names may differ, e.g. after selecting different columns per + # estimator; only shapes decide whether estimators share a call. + model = _ClassFrequencyModel() tables = [ TableTensor( columns={Stype.numerical: ("a", "b")}, - numerical=torch.ones(3, 2), + numerical=torch.randn(3, 2), ), TableTensor( columns={Stype.numerical: columns}, - numerical=torch.ones(3, len(columns)), + numerical=torch.randn(3, len(columns)), ), ] x = EnsembleTable.from_tables(tables=tables, member_table_ids=(0, 1)) - y = torch.zeros(2, 3, 1) + y = TableTensor( + categorical=CategoricalTensor( + code=torch.tensor([[0], [1], [0]]), + categories=(torch.tensor([10, 20]),), + ), + ) x_query = EnsembleTable.from_tables( tables=[table[:1] for table in tables], member_table_ids=(0, 1), ) - with pytest.raises(ValueError, match="Estimators in one batch"): - model(x, y, x_query, estimator_batch_size=None) - with pytest.raises(ValueError, match="Estimators in one batch"): - model.fit(x, y, estimator_batch_size=None) + expected = model(x, y, x_query, num_estimators=2, estimator_batch_size=1) + model.num_calls = 0 + out = model(x, y, x_query, num_estimators=2, estimator_batch_size=None) + torch.testing.assert_close(out.numerical, expected.numerical) + assert model.num_calls == num_calls -def test_estimator_batching_incompatible_class_values() -> None: + model.fit(x, y, num_estimators=2, estimator_batch_size=None) + torch.testing.assert_close( + model.predict(x_query).numerical, expected.numerical + ) + + +def test_estimator_batching_fails_like_sequential_on_class_mismatch() -> None: model = _ClassFrequencyModel() x = torch.randn(3, 2) y = EnsembleTable.from_tables( @@ -848,10 +891,97 @@ def test_estimator_batching_incompatible_class_values() -> None: ], member_table_ids=(0, 1), ) - with pytest.raises(ValueError, match="Estimators in one batch"): - model(x, y, x, num_estimators=2, estimator_batch_size=None) - with pytest.raises(ValueError, match="Estimators in one batch"): - model.fit(x, y, num_estimators=2, estimator_batch_size=None) + for estimator_batch_size in (1, None): + with pytest.raises(ValueError, match="same set of classes"): + model( + x, + y, + x, + num_estimators=2, + estimator_batch_size=estimator_batch_size, + ) + + +def test_estimator_batching_preserves_member_order() -> None: + model = _RecordingModel() + rows = (3, 4, 3, 3) + x = EnsembleTable.from_tables( + tables=[TableTensor(numerical=torch.randn(r, 2)) for r in rows], + member_table_ids=range(len(rows)), + ) + y = EnsembleTable.from_tables( + tables=[TableTensor(numerical=torch.zeros(r, 1)) for r in rows], + member_table_ids=range(len(rows)), + ) + x_query = TableTensor(numerical=torch.randn(len(rows), 2, 2)) + + out = model(x, y, x_query, estimator_batch_size=None) + + torch.testing.assert_close(out.numerical, x_query.numerical) + # Only consecutive estimators share a call: 0 | 1 | 2 and 3. + assert [ + cast(TableTensor, call.x_context).size()[:-2] for call in model.calls + ] == [(), (), (2,)] + + +def test_auto_estimator_batching_keeps_cell_budget( + monkeypatch: pytest.MonkeyPatch, +) -> None: + # Every estimator costs 3 context and 2 query rows of 2 columns plus 1. + monkeypatch.setattr(_RecordingModel, "_estimator_batch_cells", 30) + monkeypatch.setattr(_RecordingModel, "_estimator_row_cells", 1) + model = _RecordingModel() + x = torch.randn(4, 3, 2) + y = torch.zeros(4, 3, 1) + x_query = torch.randn(4, 2, 2) + + out = model(x, y, x_query) + + torch.testing.assert_close(out.numerical, x_query) + assert [ + cast(TableTensor, call.x_query).size() for call in model.calls + ] == [(2, 2, 2), (2, 2, 2)] + + +def test_auto_estimator_batching_counts_rows_without_columns( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(_RecordingModel, "_estimator_batch_cells", 9) + monkeypatch.setattr(_RecordingModel, "_estimator_row_cells", 1) + model = _RecordingModel() + + model(torch.randn(2, 3, 0), torch.zeros(2, 3, 1), torch.randn(2, 2, 0)) + + assert len(model.calls) == 2 + + +def test_auto_estimator_batching_is_sequential_with_gradients() -> None: + model = _RecordingModel() + model.train() + + model(torch.randn(3, 3, 2), torch.zeros(3, 3, 1), torch.randn(3, 2, 2)) + + assert len(model.calls) == 3 + + +def test_predict_splits_batched_query_rows_within_budget( + monkeypatch: pytest.MonkeyPatch, +) -> None: + # Two estimators with 3 context rows of 2 columns plus 1 each fit one + # batch, but their 5 query rows exceed the budget together. + monkeypatch.setattr(_RecordingModel, "_estimator_batch_cells", 18) + monkeypatch.setattr(_RecordingModel, "_estimator_row_cells", 1) + model = _RecordingModel() + model.fit(torch.randn(2, 3, 2), torch.zeros(2, 3, 1)) + model.calls.clear() + x_query = torch.randn(2, 5, 2) + + out = model.predict(x_query) + + torch.testing.assert_close(out.numerical, x_query) + assert [ + cast(TableTensor, call.x_query).size() for call in model.calls + ] == [(2, 3, 2), (2, 2, 2)] @pytest.mark.parametrize("estimator_batch_size", [1, 2, None]) @@ -880,3 +1010,81 @@ def test_forward_batching_validates_each_query_schema( ) with pytest.raises(ValueError, match="share the same schema"): model(x, y, query, estimator_batch_size=estimator_batch_size) + + +def test_predict_keeps_queries_whole_with_callbacks( + monkeypatch: pytest.MonkeyPatch, +) -> None: + # Without callbacks, this budget splits the query rows into two calls. + monkeypatch.setattr(_RecordingModel, "_estimator_batch_cells", 18) + monkeypatch.setattr(_RecordingModel, "_estimator_row_cells", 1) + model = _RecordingModel() + model.fit(torch.randn(2, 3, 2), torch.zeros(2, 3, 1)) + model.calls.clear() + events: list[str] = [] + + model.predict( + torch.randn(2, 5, 2), + callbacks=(MyCallback("affine", 2.0, 3.0, events),), + ) + + assert len(model.calls) == 1 + assert events.count("affine_model_forward_end") == 2 + + +@pytest.mark.parametrize("change", ["category_counts", "dtype"]) +def test_estimator_batching_splits_different_feature_blocks( + change: str, +) -> None: + def table(i: int) -> TableTensor: + if change == "dtype": + dtype = (torch.float32, torch.float64)[i] + return TableTensor(numerical=torch.ones(3, 2, dtype=dtype)) + return TableTensor( + categorical=CategoricalTensor( + code=torch.zeros(3, 1, dtype=torch.long), + categories=(torch.arange(2 + i),), + ), + ) + + model = _RecordingModel() + x = EnsembleTable.from_tables( + tables=[table(0), table(1)], + member_table_ids=(0, 1), + ) + + model( + x_context=x, + y_context=torch.zeros(2, 3, 1), + x_query=x, + estimator_batch_size=None, + ) + + assert len(model.calls) == 2 + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") +def test_forward_members_moves_contexts_to_query_device() -> None: + model = _RecordingModel() + recipe = RecipeExecution(model.default_recipe()) + contexts = recipe.fit_transform( + x=torch.randn(4, 3, 2), + y=torch.zeros(4, 3, 1), + related_tables=None, + num_members=None, + ) + x_query = torch.randn(4, 2, 2) + queries = [ + query._replace(x=cast(TableTensor, query.x.cuda())) + for query in recipe.transform(x=x_query, related_tables=None) + ] + + outs = model._forward_members(contexts=contexts, queries=queries) + + assert all( + cast(TableTensor, call.x_context).is_cuda for call in model.calls + ) + torch.testing.assert_close( + torch.stack([out.numerical.cpu() for out in outs]), + x_query, + ) diff --git a/test/models/test_ecoc.py b/test/models/test_ecoc.py index 30f4744ec..8563fd806 100644 --- a/test/models/test_ecoc.py +++ b/test/models/test_ecoc.py @@ -96,3 +96,89 @@ def test_ecoc( num_classes=num_classes + 1, cache=cache, ) + + +@withCUDA +def test_ecoc_members_draw_codebooks_in_order(device: torch.device) -> None: + model = MyModel(10) + ecoc = ECOC(max_classes=10) + num_members, num_classes, num_context = 3, 12, 11 + x = torch.randn( + num_members, + num_context + 2, + num_context, + device=device, + dtype=torch.float64, + ) + y = torch.rand(num_members, num_classes, device=device).argsort(dim=-1)[ + ..., :num_context + ] + + generator = torch.Generator(device).manual_seed(0) + expected = torch.stack( + [ + ecoc( + model=model, + x=x[member], + y=y[member], + num_classes=num_classes, + generator=generator, + ) + for member in range(num_members) + ] + ) + + generator = torch.Generator(device).manual_seed(0) + actual = ecoc( + model=model, + x=x, + y=y, + num_classes=num_classes, + num_members=num_members, + generator=generator, + ) + torch.testing.assert_close(actual, expected) + + cache = Cache() + ecoc( + model=model, + x=x[..., :num_context, :], + y=y, + num_classes=num_classes, + num_members=num_members, + cache=cache, + generator=torch.Generator(device).manual_seed(0), + ) + cache.freeze() + replayed = ecoc( + model=model, + x=x[..., num_context:, :], + y=y[..., :0], + num_classes=num_classes, + num_members=num_members, + cache=cache, + ) + torch.testing.assert_close(replayed, expected) + + +@pytest.mark.parametrize("num_classes", [3, 10, 11, 100, 201]) +def test_ecoc_num_tasks(num_classes: int) -> None: + ecoc = ECOC(max_classes=10) + x = torch.randn(num_classes + 2, num_classes, dtype=torch.float64) + y = torch.arange(num_classes) + cache = Cache() + + ecoc( + model=MyModel(10), + x=x[:num_classes], + y=y, + num_classes=num_classes, + cache=cache, + ) + + num_tasks = ( + cast(Tensor, cache["ecoc_codebook"]).size(-2) + if num_classes > 10 + else 1 + ) + assert ecoc.num_tasks(num_classes) == num_tasks diff --git a/test/nn/test_attention.py b/test/nn/test_attention.py index 2dee342b4..23eb0e865 100644 --- a/test/nn/test_attention.py +++ b/test/nn/test_attention.py @@ -657,7 +657,10 @@ def test_transformer_block(device: torch.device, qassmax: bool) -> None: torch.testing.assert_close(out1, out3) -def test_transformer_block_chunked_noncontiguous_out() -> None: +@pytest.mark.parametrize("dtype", [torch.float32, torch.float64]) +def test_transformer_block_chunked_noncontiguous_out( + dtype: torch.dtype, +) -> None: channels = 8 module = TransformerBlock( channels=channels, @@ -667,7 +670,8 @@ def test_transformer_block_chunked_noncontiguous_out() -> None: base = torch.randn(2, 3, 4, channels) query = base.transpose(-2, -3) - buffer = torch.empty_like(base).transpose(-2, -3) + # Under autocast, outputs can differ in dtype from a preallocated buffer. + buffer = torch.empty_like(base, dtype=dtype).transpose(-2, -3) assert not query.is_contiguous() assert not buffer.is_contiguous() @@ -681,7 +685,7 @@ def test_transformer_block_chunked_noncontiguous_out() -> None: ) assert actual is buffer - torch.testing.assert_close(actual, expected) + torch.testing.assert_close(actual, expected.to(dtype)) def test_transformer_block_kv_cache() -> None: