Source code for boundlab.ops.sparse_poly

"""Backend-dispatched sparse polynomial generator operations."""

from __future__ import annotations

import torch

from . import ref, triton
from .sum import sum

from boundlab.ibp import Noise
from boundlab.polysp.generator import Generator


[docs] def final_matmul( gen1: "Generator", gen2: "Generator", eps2: torch.Tensor, order: int, symmetric: bool, ) -> tuple["Generator", "Noise"]: """Dispatch the sparse generator pair product by device platform.""" if triton.is_supported(gen1, gen2, symmetric): return triton.final_matmul(gen1, gen2, eps2, order, symmetric) return ref.final_matmul(gen1, gen2, eps2, order, symmetric)
[docs] def matmul( gen1: "Generator", gen2: "Generator", eps: float, order: int ) -> tuple["Generator", "Noise"]: """Dispatch sparse generator matrix multiplication by device platform.""" return ref.matmul(gen1, gen2, eps=eps, order=order, final=final_matmul)
[docs] def final_quadratic( gen: "Generator", a: torch.Tensor, b: torch.Tensor, eps2: torch.Tensor, order: int, ) -> tuple["Generator", "Noise"]: """Dispatch the sparse generator quadratic pair product by device platform.""" if triton.quadratic_is_supported(gen, a, b): return triton.final_quadratic(gen, a, b, eps2, order) return ref.final_quadratic(gen, a, b, eps2, order)
[docs] def quadratic(gen: "Generator", a: torch.Tensor, b: torch.Tensor, eps: float, order: int) -> tuple["Generator", "Noise"]: """Dispatch the sparse generator quadratic form by device platform.""" return ref.quadratic(gen, a, b, eps=eps, order=order, final=final_quadratic)
__all__ = ["final_matmul", "final_quadratic", "matmul", "quadratic", "sum"]