/usr/local/lib64/python3.6/site-packages/torch/include/caffe2/perfkernels
NameSizeModeActions
adagrad.h60680644editdlrm
common.h52630644editdlrm
cvtsh_ss_bugfix.h20250644editdlrm
embedding_lookup.h15390644editdlrm
embedding_lookup_idx.h16740644editdlrm
fused_8bit_rowwise_embedding_lookup.h16900644editdlrm
fused_8bit_rowwise_embedding_lookup_idx.h18070644editdlrm
fused_nbit_rowwise_conversion.h7580644editdlrm
lstm_unit_cpu-impl.h42590644editdlrm
lstm_unit_cpu.h14460644editdlrm
lstm_unit_cpu_common.h14420644editdlrm
math.h11030644editdlrm
typed_axpy.h3550644editdlrm
Edit: /usr/local/lib64/python3.6/site-packages/torch/include/caffe2/perfkernels/embedding_lookup_idx.h (1674B)
#pragma once #include namespace caffe2 { // clang-format off /** * Embedding lookup with reduction. * * `input` of size data_size * block_size * `indices` of size index_size * `offsets` of size output_size * `weights` nullptr or array of size index_size * `out` of size output_size * block_size * * Behavior is roughly equivalent to pseudocode: * * pos = 0 * for (i = 0..output_size-1) * for (k = 0..block_size-1) * out[i*block_size + k] = 0 * start_offset = offsets[i] * end_offset = offsets[i+1] * length = end_offset - start_offset * for (j = start_offset..end_offset-1) * for (k = 0..block_size-1) * out[i*block_size + k] += input[indices[pos]*block_size + k] * * (weights ? weights[IS_WEIGHT_POSITIONAL ? j - start_offset : pos] : 1.0) * pos += 1 * if (normalize_weights && length > 0) * for (k = 0..block_size-1) * out[i*block_size + k] /= length * * TODO: make this API also take "offsets" rather than "lengths" to match the * API for PyTorch's EmbeddingBag */ // clang-format on template < typename IndexType, typename InType, typename OutType, bool IS_WEIGHT_POSITIONAL = false> void EmbeddingLookupIdx( const std::int64_t block_size, const std::int64_t output_size, const std::int64_t index_size, const std::int64_t data_size, const InType* input, const IndexType* indices, const IndexType* offsets, const float* weights, // optional, can be null for non-weighted sum const float* scale_bias, // optional scale & bias params for uint8 input bool normalize_by_lengths, OutType* out); } // namespace caffe2