Contents

namespace

SushiBLAS::Dispatch

Declared in
include/SushiBLAS/engine/blas/argument_check.hpp

Contains

Enumerations

enum class Domain

Enumerates the element-type families a dispatch helper may instantiate its kernel for.

Warning

Widen it from REAL only when the functor is well-formed for the family being added.

See also

include/SushiBLAS/engine/README.md

REAL

HALF, FLOAT32, FLOAT64.

COMPLEX

COMPLEX32, COMPLEX64.

INTEGRAL

INT32, INT64.

REAL_AND_COMPLEX

Everything but the integers.

REAL_AND_INTEGRAL

Everything an ordering is defined on.

ALL

Every element type.

enum class Sampling

Names what a random op does with element values, which decides whether it accepts HALF.

See also

include/SushiBLAS/engine/math/README.md

DRAWS_VALUES

Computes new values; HALF has too little mantissa and is refused.

MOVES_VALUES

Fills or permutes values; every dtype of its domain is accepted.

Variables

constexpr int K_RNG_SEED_PARAM = 6

Holds the task-metadata slot of the RNG seed, stored as raw uint64_t bits.

See also

include/SushiBLAS/engine/math/README.md

constexpr int K_RNG_OFFSET_PARAM = 7

Task-metadata slot holding the per-task RNG offset, as raw uint64_t bits.

constexpr size_t K_MAX_RNG_PARAMS = 6

Holds how many distribution parameters a random op may record: the slots left once K_RNG_SEED_PARAM and K_RNG_OFFSET_PARAM are reserved.

Functions

template <typename Flag>
void require_flag(Flag flag, const char *op, const char *argument)

Throws unless a flag holds one of its type's enumerators.

Template parameters

Flag

Core::Transpose, Core::Uplo, Core::Diag or Core::Side.

Exceptions

std::invalid_argument

naming the argument and its numeric value.

void require_real_scalar(const BlasScalar &value, Core::DataType dtype, const char *op, const char *argument)

Throws if a scalar with a non-zero imaginary part is given to a real dtype.

Exceptions

std::invalid_argument

naming the imaginary part and the dtype.

void require_real_dtype(const Tensor &t, const char *op, const char *operand, const char *siblings)

Throws if a routine defined for real tensors is given a complex one.

Parameters

siblings

Names the routines that take a complex tensor, such as "dotu or dotc".

Exceptions

std::invalid_argument

naming the dtype and the siblings.

void require_transpose_not(Core::Transpose trans, Core::Transpose refused, const char *op, const char *sibling)

Throws if trans is the one value a routine does not define for a complex tensor.

Parameters

sibling

Names the routine that takes the refused value.

Exceptions

std::invalid_argument

naming the value and the sibling.

template <typename Epilogue>
GemmEpilogue< Epilogue > gemm_epilogue(Epilogue functor)

Bundles functor into a GemmEpilogue that declares nothing yet.

See also

include/SushiBLAS/engine/blas/README.md

template <Domain Dom = Domain::REAL, typename Epilogue>
TaskHandle execute_gemm(TaskRecorder &recorder, const Tensor &A, const Tensor &B, Tensor &C, bool transA, bool transB, double alpha, double beta, const char *name, SushiRuntime::Graph::OpID op_id, Epilogue epilogue)

Records C = epilogue(alpha * op(A) * op(B) + beta * C, row, col) as one task.

Parameters

C

Must not overlap A or B.

epilogue

Must be device-copyable, capture by value, and neither allocate nor throw.

Returns

The handle of the recorded task; the handle of no task when nothing was recorded.

Exceptions

std::runtime_error

on rank < 2, a dtype mismatch, or a dtype outside Dom.

std::invalid_argument

on a non-contiguous operand, mixed layouts, or a shape or batch mismatch.

See also

include/SushiBLAS/engine/blas/README.md

MatrixOperand describe_matrix_operand(const Tensor &t, const char *op, const char *operand)

Returns the geometry of a dense matrix operand of rank 2 or more.

Exceptions

std::invalid_argument

on rank < 2, missing storage, a non-contiguous tensor or a dtype the tensor's device cannot compute in.

See also

include/SushiBLAS/engine/blas/README.md

int64_t batch_stride(const MatrixOperand &m, int64_t batch_count, const char *op, const char *operand)

Returns the element distance between consecutive matrices of operand within a call.

Parameters

batch_count

The number of matrices the call computes.

Returns

0 when operand has rank 2 or a batch count of 1, matrix_elements otherwise.

Exceptions

std::invalid_argument

when operand is batched with a different batch count.

void require_same_batch_axes(const Tensor &a, const MatrixOperand &ma, const Tensor &b, const MatrixOperand &mb, const char *op, const char *operand_a, const char *operand_b)

Throws unless two operands that both hold more than one matrix share their batch axes.

Precondition

a and b have passed describe_matrix_operand, which gave ma and mb.

Exceptions

std::invalid_argument

when the axes before the last two differ in count or size.

void require_same_layout(const MatrixOperand &a, const MatrixOperand &b, const char *op, const char *operand_a, const char *operand_b)

Throws unless two matrix operands share one layout.

Exceptions

std::invalid_argument

naming both layouts on a mismatch.

void require_matrix_extent(int64_t actual, int64_t expected, const char *what, const char *op, const char *operand)

Throws unless one matrix extent of an operand equals the extent the call needs.

Parameters

what

Names the extent in the message, such as "rows of op(A)".

Exceptions

std::invalid_argument

naming both extents on a mismatch.

void require_unbatched(const MatrixOperand &m, const char *op, const char *operand)

Throws unless a matrix operand holds exactly one matrix of rank 2.

Exceptions

std::invalid_argument

naming the rank.

void require_square(const MatrixOperand &m, const char *op, const char *operand)

Throws unless a matrix operand has as many rows as columns.

Exceptions

std::invalid_argument

naming both extents.

VectorOperand describe_vector_operand(const Tensor &t, const char *op, const char *operand)

Returns the geometry of a vector operand.

Exceptions

std::invalid_argument

on missing storage, a dtype the tensor's device cannot compute in, a non-contiguous tensor of rank above 1, or a rank-1 stride of zero or less.

void require_vector_length(const VectorOperand &v, int64_t expected, const char *against, const char *op, const char *operand)

Throws unless a vector holds the number of elements the call needs.

Parameters

against

Names what sets the count, such as "the columns of op(A)".

Exceptions

std::invalid_argument

naming both counts.

void require_unit_increment(const VectorOperand &v, const char *op, const char *operand)

Throws unless a vector's elements are adjacent.

Exceptions

std::invalid_argument

naming the increment.

constexpr bool domain_admits(Domain allowed, Domain family) noexcept

Returns whether allowed includes the element-type family.

Parameters

family

A single-bit Domain naming one family.

See also

include/SushiBLAS/engine/README.md

constexpr Domain domain_of(Core::DataType dtype) noexcept

Returns the element-type family a runtime dtype belongs to.

Returns

REAL, COMPLEX or INTEGRAL, never a combination.

See also

include/SushiBLAS/engine/README.md

template <Domain Dom>
void require_domain(Core::DataType dtype, const char *family, const char *name)

Throws unless dtype is one Dom admits.

Exceptions

std::runtime_error

naming family, name and dtype when the dtype falls outside Dom.

See also

include/SushiBLAS/engine/README.md

template <Domain Dom, typename Arm>
void with_element_type(Core::DataType dtype, const char *family, const char *name, Arm &&arm)

Runs arm instantiated for the C++ type dtype denotes, or throws.

Parameters

arm

Callable as arm(ElementTag<T>{}) or arm.template operator()<T>() for each type in Dom.

Exceptions

std::runtime_error

naming family, name and dtype when the dtype falls outside Dom.

See also

include/SushiBLAS/engine/README.md

template <typename T>
void set_task_params(SushiRuntime::Graph::TaskMetadata &meta, const std::vector< T > &params, const char *name, size_t reserved_slots=0)

Copies a parameter list into a task's inline metadata slots, starting at slot 0.

Template parameters

T

The parameter element type; must fit in 64 bits.

Parameters

reserved_slots

Trailing slots the caller keeps for itself, which params must not reach.

See also

include/SushiBLAS/engine/README.md

template <Domain Dom = Domain::REAL, typename Func, typename Param = double>
TaskHandle execute_logic_unary(TaskRecorder &recorder, const Tensor &x, Tensor &result, const char *name, SushiRuntime::Graph::OpID op_id, Func &&op_func, const std::vector< Param > &params={})

Records a unary logic op result[i] = op(x[i]), or nothing when result is empty.

Parameters

result

Must match x in size and dtype.

op_func

Callable (auto x) -> auto: device-copyable, by-value captures, no allocation, no throw.

Exceptions

std::invalid_argument

on a missing storage, a non-contiguous operand, a shape or layout mismatch, or under FORCE_REFERENCE.

std::runtime_error

on a dtype mismatch, a dtype outside Dom, or too many params.

See also

include/SushiBLAS/engine/logic/README.md

template <Domain Dom = Domain::REAL, typename Func, typename Param = double>
TaskHandle execute_logic_binary(TaskRecorder &recorder, const Tensor &A, const Tensor &B, Tensor &result, const char *name, SushiRuntime::Graph::OpID op_id, Func &&op_func, const std::vector< Param > &params={})

Records a binary logic op result[i] = op(A[i], B[i]), or nothing when result is empty.

Template parameters

Dom

Defaults to Domain::REAL, because < and > have no meaning on complex numbers.

Parameters

result

Must match the inputs in size and dtype.

op_func

Callable (auto x, auto y) -> auto under execute_logic_unary's constraints.

Exceptions

std::invalid_argument

on a missing storage, a non-contiguous operand, a shape or layout mismatch, or under FORCE_REFERENCE.

std::runtime_error

on a dtype mismatch, a dtype outside Dom, or too many params.

template <Domain Dom = Domain::REAL, typename Func, typename Param = double>
TaskHandle execute_logic_ternary(TaskRecorder &recorder, const Tensor &cond, const Tensor &A, const Tensor &B, Tensor &result, const char *name, SushiRuntime::Graph::OpID op_id, Func &&op_func, const std::vector< Param > &params={})

Records a ternary logic op result[i] = op(cond[i], A[i], B[i]), or nothing when result is empty.

Parameters

cond

Nonzero selects A, zero selects B.

result

Must match the inputs in size and dtype.

op_func

Callable (auto c, auto a, auto b) -> auto under execute_logic_unary's constraints.

Exceptions

std::invalid_argument

on a missing storage, a non-contiguous operand, a shape or layout mismatch, or under FORCE_REFERENCE.

std::runtime_error

on a dtype mismatch, a dtype outside Dom, or too many params.

template <typename T>
bool logic_not_zero(T x)

Truth test for a real scalar.

Parameters

x

The value to test.

Returns

True if x is not zero.

template <typename T>
bool logic_not_zero(const std::complex< T > &x)

Truth test for a complex value.

Parameters

x

The value to test.

Returns

True if either component of x is nonzero.

template <typename T>
bool logic_equal(T x, T y)

Equality of two real scalars.

Parameters

x

Left operand.

y

Right operand.

Returns

True if x == y.

template <typename T>
bool logic_equal(const std::complex< T > &x, const std::complex< T > &y)

Equality of two complex values, component-wise.

Parameters

x

Left operand.

y

Right operand.

Returns

True if both components match.

template <typename T>
bool logic_less(T x, T y)

Less-than comparison of two real scalars.

Parameters

x

Left operand.

y

Right operand.

Returns

True if x < y.

template <typename T>
bool logic_greater(T x, T y)

Greater-than comparison of two real scalars.

Parameters

x

Left operand.

y

Right operand.

Returns

True if x > y.

template <typename T>
bool logic_less_equal(T x, T y)

Less-than-or-equal comparison of two real scalars.

Parameters

x

Left operand.

y

Right operand.

Returns

True if x <= y.

template <typename T>
bool logic_greater_equal(T x, T y)

Greater-than-or-equal comparison of two real scalars.

Parameters

x

Left operand.

y

Right operand.

Returns

True if x >= y.

template <Domain Dom = Domain::REAL, typename Func, typename Param = double>
TaskHandle execute_unary_inplace(TaskRecorder &recorder, Tensor &t, const char *name, SushiRuntime::Graph::OpID op_id, Func &&op_func, const std::vector< Param > &params={})

Records a unary in-place elementwise op t[i] = op(t[i]), or nothing when t is empty.

Template parameters

Dom

Defaults to Domain::REAL; REAL_AND_COMPLEX needs an op_func well-formed for it.

Parameters

op_func

Callable (auto x) -> auto: device-copyable, by-value captures, no allocation, no throw.

Exceptions

std::invalid_argument

if t holds elements without storage or is not contiguous, or under FORCE_REFERENCE.

std::runtime_error

if the dtype falls outside Dom or params exceeds the inline capacity.

See also

include/SushiBLAS/engine/math/README.md

template <Domain Dom = Domain::REAL, typename Func, typename Param = double>
TaskHandle execute_binary(TaskRecorder &recorder, const Tensor &A, const Tensor &B, Tensor &C, const char *name, SushiRuntime::Graph::OpID op_id, Func &&op_func, const std::vector< Param > &params={})

Records a binary elementwise op C[i] = op(A[i], B[i]), or nothing when C is empty.

Template parameters

Dom

Defaults to Domain::REAL; REAL_AND_COMPLEX needs an op_func well-formed for it.

Parameters

C

May alias an input.

op_func

Callable (auto x, auto y) -> auto under execute_unary_inplace's constraints.

Exceptions

std::invalid_argument

on a missing storage, a non-contiguous operand, a shape or layout mismatch, or under FORCE_REFERENCE.

std::runtime_error

on a dtype mismatch, a dtype outside Dom, or too many params.

See also

include/SushiBLAS/engine/math/README.md

template <Domain Dom = Domain::REAL, std::size_t N, typename Func, typename Param = double>
TaskHandle execute_zip_inplace(TaskRecorder &recorder, const std::array< Tensor *, N > &tensors, const char *name, SushiRuntime::Graph::OpID op_id, Func &&op_func, const std::vector< Param > &params={})

Records a fused in-place op over N same-shaped tensors, or nothing when they are empty.

Parameters

tensors

None may be null; all must share one shape, layout and dtype, and none may overlap.

op_func

Callable (auto&...) -> void taking N lvalue references.

Exceptions

std::invalid_argument

on a missing storage, a non-contiguous operand, a shape or layout mismatch, or under FORCE_REFERENCE.

std::runtime_error

on a null tensor, a dtype mismatch, or a dtype outside Dom.

See also

include/SushiBLAS/engine/math/README.md

template <typename T>
T safe_ew_add(T a, T b)

Returns a + b, wrapping modulo 2^N when T is an integer type.

See also

include/SushiBLAS/engine/math/README.md

template <typename T>
T safe_ew_sub(T a, T b)

Returns a - b, wrapping modulo 2^N when T is an integer type.

See also

include/SushiBLAS/engine/math/README.md

template <typename T>
T safe_ew_mul(T a, T b)

Returns a * b, wrapping modulo 2^N when T is an integer type.

See also

include/SushiBLAS/engine/math/README.md

template <typename T>
T safe_ew_div(T a, T b)

Division, as a named functor for callers composing ops generically.

Parameters

a

Left operand.

b

Right operand.

Returns

a / b.

template <typename Fn, typename... Scalars>
ScalarBound< Fn, sizeof...(Scalars)> bind_scalars(Fn fn, Scalars... scalars)

Binds scalar parameters to an element functor for a seam to narrow per element type.

Parameters

fn

Callable (elements..., auto scalars...); it takes the scalars after the elements.

See also

include/SushiBLAS/engine/math/README.md

template <typename T>
T wrapping_add(T a, T b)

Returns a + b, wrapping modulo 2^N when T is an integer type.

See also

include/SushiBLAS/engine/math/README.md

template <typename T>
T wrapping_sub(T a, T b)

Returns a - b, wrapping modulo 2^N when T is an integer type.

See also

include/SushiBLAS/engine/math/README.md

template <typename T>
T wrapping_mul(T a, T b)

Returns a * b, wrapping modulo 2^N when T is an integer type.

See also

include/SushiBLAS/engine/math/README.md

template <typename T>
T wrapping_neg(T x)

Returns -x, wrapping modulo 2^N when T is an integer type.

See also

include/SushiBLAS/engine/math/README.md

template <typename T>
T wrapping_abs(T x)

Returns |x|; the most negative integer maps to itself.

See also

include/SushiBLAS/engine/math/README.md

template <Domain Dom = Domain::REAL, typename Func, typename Param = double>
TaskHandle execute_nonlinear_forward(TaskRecorder &recorder, Tensor &t, const char *name, SushiRuntime::Graph::OpID op_id, Func &&op_func, const std::vector< Param > &params={})

Records an activation's forward pass t[i] = op(t[i]), or nothing when t is empty.

Template parameters

Dom

Defaults to Domain::REAL; REAL_AND_COMPLEX needs an op_func well-formed for it.

Parameters

op_func

Callable (auto x) -> auto: device-copyable, by-value captures, no allocation, no throw.

Exceptions

std::invalid_argument

if t holds elements without storage or is not contiguous, or under FORCE_REFERENCE.

std::runtime_error

if the dtype falls outside Dom or params exceeds the inline capacity.

See also

include/SushiBLAS/engine/math/README.md

template <Domain Dom = Domain::REAL, typename Func, typename Param = double>
TaskHandle execute_nonlinear_backward(TaskRecorder &recorder, const Tensor &dy, const Tensor &x, Tensor &dx, const char *name, SushiRuntime::Graph::OpID op_id, Func &&op_func, const std::vector< Param > &params={})

Records an activation's backward pass dx[i] = op(dy[i], x[i]), or nothing when dx is empty.

Parameters

x

The saved forward input or output the derivative is taken against.

op_func

Callable (auto dy, auto x) -> auto under execute_nonlinear_forward's constraints.

Exceptions

std::invalid_argument

on a missing storage, a non-contiguous operand, a shape or layout mismatch, or under FORCE_REFERENCE.

std::runtime_error

on a dtype mismatch, a dtype outside Dom, or too many params.

See also

include/SushiBLAS/engine/math/README.md

void set_rng_params(SushiRuntime::Graph::TaskMetadata &meta, uint64_t seed, uint64_t offset)

Stamps the seed and offset a random task draws from into its metadata.

constexpr bool random_admits(Domain domain, Sampling sampling, Core::DataType dtype) noexcept

Reports whether a random op of domain and sampling accepts dtype.

See also

include/SushiBLAS/engine/math/README.md

void require_random_dtype(const Tensor &t, Domain domain, Sampling sampling, const char *name)

Throws std::invalid_argument naming the op and t's dtype unless random_admits holds.

See also

include/SushiBLAS/engine/math/README.md

template <typename T, typename DrawBlock>
sycl::event submit_random_fill(sycl::queue &q, const std::vector< sycl::event > &deps, int64_t n, T *out, DrawBlock draw_block)

Submits one counter-based fill in blocks: out[b*K + k] = draw_block(b)[k] below n.

Parameters

draw_block

Callable (uint64_t block) -> std::array<T, K>, inlined into a SYCL device kernel.

See also

include/SushiBLAS/engine/math/README.md

template <Domain Dom, typename Func>
TaskHandle execute_random(TaskRecorder &recorder, RngStream &stream, Tensor &t, const char *name, SushiRuntime::Graph::OpID op_id, const std::vector< double > &params, Func &&task_func)

Records a random-fill op and resolves the tensor's dtype to a typed pointer.

Template parameters

Dom

The element families the distribution draws; HALF is always refused.

Parameters

params

The distribution's parameters, recorded in metadata slots 0..n-1.

task_func

Submits the fill for one element type of Dom; whatever it submits runs on the device.

Exceptions

std::invalid_argument

for a dtype outside Dom, HALF, no storage, a non-contiguous t, or under FORCE_REFERENCE.

See also

include/SushiBLAS/engine/math/README.md

std::size_t partial_bytes(Core::DataType dtype)

Returns the bytes one partial of a floating-point reduction over dtype takes.

Returns

sizeof(Reduce::AccumulationOf<T>) for the element type T of dtype.

std::size_t planned_partials_for(const Tensor &t)

Returns the partials a whole-tensor reduction over t plans on t's own device.

Returns

The partial count, or 0 when t holds no elements.

Precondition

require_storage(t) has passed.

See also

include/SushiBLAS/engine/math/README.md

template <typename Dispatcher>
TaskHandle execute_reduction(TaskRecorder &recorder, const Tensor &t, Tensor &result, const char *name, SushiRuntime::Graph::OpID op_id, Dispatcher &&dispatch, std::size_t scratch_elements=0)

Records a whole-tensor reduction to a scalar result tensor.

Parameters

result

Must hold one element and share t's dtype.

dispatch

Builds the command group per element type and must call depends_on(deps).

scratch_elements

Per-work-group partials to allocate; zero allocates nothing.

Exceptions

std::invalid_argument

if an operand holds elements without storage, t is not contiguous, or under FORCE_REFERENCE.

std::runtime_error

if result is not scalar, the dtypes disagree, or the dtype is not real.

See also

include/SushiBLAS/engine/math/README.md

template <typename Dispatcher>
TaskHandle execute_index_reduction(TaskRecorder &recorder, const Tensor &t, Tensor &result, const char *name, SushiRuntime::Graph::OpID op_id, Dispatcher &&dispatch)

Records a whole-tensor reduction whose result is an index in an INT64 scalar.

Parameters

result

Must be a scalar INT64 tensor.

dispatch

Builds the command group per input element type, as in execute_reduction.

Exceptions

std::invalid_argument

if an operand holds elements without storage, t is not contiguous, or under FORCE_REFERENCE.

std::runtime_error

if result is not an INT64 scalar or the input dtype is complex.

See also

include/SushiBLAS/engine/math/README.md

template <typename Dispatcher>
TaskHandle execute_counting_reduction(TaskRecorder &recorder, const Tensor &t, Tensor &result, const char *name, SushiRuntime::Graph::OpID op_id, Dispatcher &&dispatch, std::size_t scratch_elements=0)

Records a whole-tensor reduction that counts elements into an INT64 scalar.

Parameters

result

Must be a scalar INT64 tensor.

scratch_elements

Per-work-group partials to allocate; zero allocates nothing.

Exceptions

std::invalid_argument

if an operand holds elements without storage, t is not contiguous, or under FORCE_REFERENCE.

std::runtime_error

if result is not an INT64 scalar or the input dtype is not real.

See also

include/SushiBLAS/engine/math/README.md

template <template< typename > class Accum, typename Transform, typename Finalize>
TaskHandle execute_folding_reduction(TaskRecorder &recorder, const Tensor &t, Tensor &result, const char *name, SushiRuntime::Graph::OpID op_id, Transform transform, Finalize finalize)

Records a whole-tensor reduction as result = finalize(fold(transform(t[i])), n).

Template parameters

Accum

Reduce::CompensatedSum for anything additive, Reduce::RunningProduct for a product.

Parameters

transform

Callable (auto x) -> auto, applied to each element before the fold.

finalize

Callable (auto value, std::size_t n) -> auto, applied once to the folded result.

Exceptions

std::invalid_argument

for the same reasons execute_reduction does.

std::runtime_error

for the same reasons execute_reduction does.

See also

include/SushiBLAS/engine/math/README.md

Reduce::AxisSpan axis_span_for(const Tensor &t, int axis, const Tensor &result, const char *name)

Resolves a rank-2 tensor and an axis into the stride pair a segmented reduction walks.

Parameters

axis

0 reduces down the columns, 1 along the rows.

result

Rank-1 tensor with one element per surviving position.

Exceptions

std::runtime_error

on a rank, axis or length mismatch.

See also

include/SushiBLAS/engine/math/README.md

Reduce::AxisSpan spatial_span_for(const Tensor &t, const Tensor &result, const char *name)

Resolves an NHWC tensor into the geometry that folds H and W away: [N, H, W, C] -> [N, C].

Parameters

t

Must be rank 4, NHWC, with rows contiguous with one another.

result

Rank-2 [N, C] tensor with one element per surviving pair.

Exceptions

std::runtime_error

on a rank or shape mismatch, or a row-sliced view.

See also

include/SushiBLAS/engine/math/README.md

template <template< typename > class Accum, typename Transform, typename Finalize>
TaskHandle execute_folding_span_reduction(TaskRecorder &recorder, const Tensor &t, const Reduce::AxisSpan &span, Tensor &result, const char *name, SushiRuntime::Graph::OpID op_id, Transform transform, Finalize finalize)

Records a segmented reduction over a span, one result per segment.

Parameters

span

The segment geometry, from axis_span_for or spatial_span_for.

result

Must share t's dtype.

finalize

Callable (auto value, std::size_t n) -> auto, applied once per segment.

Exceptions

std::invalid_argument

if an operand holds elements without storage, or under FORCE_REFERENCE.

std::runtime_error

if the dtypes disagree or the dtype is not HALF, FLOAT32 or FLOAT64.

See also

include/SushiBLAS/engine/math/README.md

template <template< typename > class Accum, typename Transform, typename Finalize>
TaskHandle execute_folding_axis_reduction(TaskRecorder &recorder, const Tensor &t, int axis, Tensor &result, const char *name, SushiRuntime::Graph::OpID op_id, Transform transform, Finalize finalize)

Records a segmented reduction over one axis of a rank-2 tensor.

Parameters

axis

Which axis to fold away; see axis_span_for.

result

Rank-1 tensor sharing t's dtype.

Exceptions

std::invalid_argument

if an operand holds elements without storage, or under FORCE_REFERENCE.

std::runtime_error

for the reasons axis_span_for and execute_folding_span_reduction list.

See also

include/SushiBLAS/engine/math/README.md

template <typename T>
T real_part(T x)

Identity for a real scalar; the counterpart to the complex overload.

Parameters

x

The value to project.

Returns

x unchanged.

template <typename T>
std::complex< T > real_part(const std::complex< T > &x)

Drops a complex value's imaginary component, keeping the complex type.

Parameters

x

The value to project.

Returns

x.real() re-wrapped as a complex number with zero imaginary part.

template <typename T>
T make_complex_safe(T val, T)

Widens a real scalar to the element type in use; identity for real types.

Parameters

val

The value to widen.

Returns

val unchanged.

template <typename T, typename U>
std::complex< T > make_complex_safe(U val, const std::complex< T > &)

Widens a real scalar into a complex element type.

Parameters

val

The real value to widen.

Returns

val as a complex number with zero imaginary part.

template <typename T>
bool greater_than_zero(T x)

Sign test for a real scalar.

Parameters

x

The value to test.

Returns

True if x > 0.

template <typename T>
bool greater_than_zero(const std::complex< T > &x)

Sign test for a complex value, taken on the real part.

Parameters

x

The value to test.

Returns

True if x.real() > 0.

template <typename T>
bool less_than_zero(T x)

Sign test for a real scalar.

Parameters

x

The value to test.

Returns

True if x < 0.

template <typename T>
bool less_than_zero(const std::complex< T > &x)

Sign test for a complex value, taken on the real part.

Parameters

x

The value to test.

Returns

True if x.real() < 0.

template <typename T>
T safe_exp(T x)

Exponential of a real scalar.

Parameters

x

The exponent.

Returns

e**x.

template <typename T>
std::complex< T > safe_exp(const std::complex< T > &x)

Exponential of a complex value.

Parameters

x

The exponent.

Returns

e**x.

template <typename T>
T safe_tanh(T x)

Hyperbolic tangent of a real scalar.

Parameters

x

The argument.

Returns

tanh(x).

template <typename T>
std::complex< T > safe_tanh(const std::complex< T > &x)

Hyperbolic tangent of a complex value.

Parameters

x

The argument.

Returns

tanh(x).

template <typename T>
T safe_erf(T x)

Error function of a real scalar.

Parameters

x

The argument.

Returns

erf(x).

template <typename T>
std::complex< T > safe_erf(const std::complex< T > &x)

Returns the error function of the real part of a complex value, with zero imaginary part.

See also

include/SushiBLAS/engine/math/README.md

template <typename T>
T safe_log(T x)

Natural logarithm of a real scalar.

Parameters

x

The argument.

Returns

log(x).

template <typename T>
std::complex< T > safe_log(const std::complex< T > &x)

Returns the principal natural logarithm of x, branch cut along the negative real axis.

See also

include/SushiBLAS/engine/math/README.md

template <typename T>
T safe_sqrt(T x)

Square root of a real scalar.

Parameters

x

The argument.

Returns

sqrt(x).

template <typename T>
std::complex< T > safe_sqrt(const std::complex< T > &x)

Returns the principal square root of x, the root with non-negative real part.

See also

include/SushiBLAS/engine/math/README.md

template <typename T>
T safe_abs(T x)

Absolute value of a real scalar.

Parameters

x

The argument.

Returns

|x|.

template <typename T>
std::complex< T > safe_abs(const std::complex< T > &x)

Returns the modulus of a complex value as (|x|, 0) in the complex element type.

See also

include/SushiBLAS/engine/math/README.md

template <typename T>
T safe_conj(T x) noexcept

Returns the complex conjugate of x; a real x is returned unchanged.

See also

include/SushiBLAS/engine/math/README.md

template <typename T>
T safe_mul(T a, T b) noexcept

Returns a * b without calling a std::complex operator.

See also

include/SushiBLAS/engine/math/README.md

template <typename T>
T safe_div(T a, T b) noexcept

Returns a / b without calling a std::complex operator.

See also

include/SushiBLAS/engine/math/README.md

void require_storage(const Tensor &t, const char *op, const char *operand)

Throws unless a tensor holding elements has storage.

Exceptions

std::invalid_argument

when num_elements > 0 and storage is null.

See also

include/SushiBLAS/engine/README.md

void require_dtype(const Tensor &t, Core::DataType expected, const char *op, const char *operand)

Throws unless a tensor holds the expected dtype.

Exceptions

std::invalid_argument

naming both dtypes on a mismatch.

See also

include/SushiBLAS/engine/README.md

void require_same_shape(const Tensor &a, const Tensor &b, const char *op, const char *operand_a, const char *operand_b)

Throws unless two tensors have the same rank and the same size in every dimension.

Exceptions

std::invalid_argument

naming the first rank or dimension that differs.

See also

include/SushiBLAS/engine/README.md

void require_same_layout(const Tensor &a, const Tensor &b, const char *op, const char *operand_a, const char *operand_b)

Throws unless two tensors of rank 2 or more share one memory layout.

Exceptions

std::invalid_argument

naming both layouts on a mismatch.

See also

include/SushiBLAS/engine/README.md

void require_same_element_count(const Tensor &a, const Tensor &b, const char *op, const char *operand_a, const char *operand_b)

Throws unless two tensors hold the same number of elements.

Exceptions

std::invalid_argument

naming both counts on a mismatch.

See also

include/SushiBLAS/engine/README.md

void require_min_elements(const Tensor &t, int64_t min_elements, const char *op, const char *operand)

Throws unless a tensor holds at least min_elements elements.

Exceptions

std::invalid_argument

naming the count and the minimum.

See also

include/SushiBLAS/engine/README.md

void require_contiguous(const Tensor &t, const char *op, const char *operand)

Throws unless a tensor's strides are the dense strides of its shape under its layout.

Exceptions

std::invalid_argument

on a zero or negative stride, a transposed view or a slice.

See also

include/SushiBLAS/engine/README.md

template <typename... Args>
void require_operand(bool ok, const char *op, const char *operand, const char *fmt, const Args &... args)

Throws unless ok holds; the form for a rule no named check states.

Parameters

fmt

Format of the violated rule, read after "operand '<operand>' ".

Exceptions

std::invalid_argument

naming the operation, the operand and the rule.

See also

include/SushiBLAS/engine/README.md

void require_rank(const Tensor &t, int32_t expected, const char *op, const char *operand)

Throws unless a tensor has the expected rank.

Exceptions

std::invalid_argument

naming both ranks.

See also

include/SushiBLAS/engine/README.md

void require_scalar(const Tensor &t, const char *op, const char *operand)

Throws unless a tensor holds exactly one element.

Exceptions

std::invalid_argument

naming the count.

See also

include/SushiBLAS/engine/README.md

void require_extent(int64_t actual, int64_t expected, const char *what, const char *op, const char *operand)

Throws unless an extent of an operand equals the one another operand fixes.

Parameters

what

Names the extent in the message, such as "channels" or "rows".

Exceptions

std::invalid_argument

naming both extents.

See also

include/SushiBLAS/engine/README.md

void require_device_dtype(const Tensor &t, const char *op, const char *operand)

Throws unless the device a tensor lives on can compute in the tensor's dtype.

Exceptions

std::invalid_argument

for FLOAT64 or COMPLEX64 storage on a device without the fp64 aspect.

See also

include/SushiBLAS/engine/README.md