"""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"]