Skip to main content

ClikaRT::ops::embedding

function

embedding()​

Tensor embedding(
    Tensor indices,
    Tensor weight,
    OptionalTensor bias = {},
    std::optional<Activation> activation = std::nullopt
)

Embedding lookup: gathers rows of weight by indices, with an optional bias + activation epilogue.

Parameters

  • indices: Int32 or Int64, any shape; every value in [0, V).
  • weight: the table [V, E] (row per vocabulary entry).
  • bias: optional [E], added to every gathered row.
  • activation: optional activation applied after the bias.

An index outside the table (negative, or at or past V) is refused, on every backend, with INVALID_ARGUMENT naming the offending value and the table size, the same contract as index_select; no index is ever clamped or wrapped to another row.

Throws

  • ClikaRT::Error: (INVALID_ARGUMENT) when indices is not an integer dtype, an index is out of range, or the shapes are inconsistent (code_name() carries the reason).

Returns: indices.shape ++ [E], dtype of weight.

auto h = ClikaRT::ops::embedding(token_ids, table); // [B, S] -> [B, S, E]

Declared in ClikaRT/compute/ops.h, line 2092