/
usr
/
local
/
lib64
/
python3.6
/
site-packages
/
torch
/
include
/
caffe2
/
operators
/
/usr/local/lib64/python3.6/site-packages/torch/include/caffe2/operators
mkdir
upload
Name
Size
Mode
Actions
abs_op.h
705
0644
edit
dl
rm
accumulate_op.h
1073
0644
edit
dl
rm
accuracy_op.h
652
0644
edit
dl
rm
acos_op.h
711
0644
edit
dl
rm
activation_ops_cudnn.h
4122
0644
edit
dl
rm
affine_channel_op.h
3450
0644
edit
dl
rm
alias_with_name.h
1234
0644
edit
dl
rm
apmeter_op.h
1027
0644
edit
dl
rm
arg_ops.h
2319
0644
edit
dl
rm
asin_op.h
711
0644
edit
dl
rm
assert_op.h
1335
0644
edit
dl
rm
async_net_barrier_op.h
904
0644
edit
dl
rm
atan_op.h
711
0644
edit
dl
rm
batch_box_cox_op.h
2287
0644
edit
dl
rm
batch_bucketize_op.h
720
0644
edit
dl
rm
batch_gather_ops.h
5264
0644
edit
dl
rm
batch_matmul_op.h
9602
0644
edit
dl
rm
batch_moments_op.h
3364
0644
edit
dl
rm
batch_permutation_op.h
954
0644
edit
dl
rm
batch_sparse_to_dense_op.h
6147
0644
edit
dl
rm
bbox_transform_op.h
2668
0644
edit
dl
rm
bisect_percentile_op.h
4921
0644
edit
dl
rm
boolean_mask_ops.h
2665
0644
edit
dl
rm
boolean_unmask_ops.h
378
0644
edit
dl
rm
box_with_nms_limit_op.h
4960
0644
edit
dl
rm
bucketize_op.h
1361
0644
edit
dl
rm
byte_weight_dequant_op.h
1722
0644
edit
dl
rm
cast_op.h
1393
0644
edit
dl
rm
cbrt_op.h
723
0644
edit
dl
rm
cc_bmm_bg_op.h
3894
0644
edit
dl
rm
ceil_op.h
782
0644
edit
dl
rm
channel_backprop_stats_op.h
737
0644
edit
dl
rm
channel_shuffle_op.h
1902
0644
edit
dl
rm
channel_stats_op.h
1807
0644
edit
dl
rm
clip_op.h
1639
0644
edit
dl
rm
collect_and_distribute_fpn_rpn_proposals_op.h
6875
0644
edit
dl
rm
concat_split_op.h
11850
0644
edit
dl
rm
conditional_op.h
487
0644
edit
dl
rm
conv_op.h
3125
0644
edit
dl
rm
conv_op_cache_cudnn.h
1935
0644
edit
dl
rm
conv_op_impl.h
28729
0644
edit
dl
rm
conv_op_shared.h
672
0644
edit
dl
rm
conv_pool_op_base.h
32109
0644
edit
dl
rm
conv_transpose_op.h
1727
0644
edit
dl
rm
conv_transpose_op_impl.h
18264
0644
edit
dl
rm
conv_transpose_op_mobile.h
1470
0644
edit
dl
rm
conv_transpose_op_mobile_impl.h
19587
0644
edit
dl
rm
conv_transpose_unpool_op_base.h
10303
0644
edit
dl
rm
copy_op.h
1296
0644
edit
dl
rm
copy_rows_to_tensor_op.h
2599
0644
edit
dl
rm
cosh_op.h
711
0644
edit
dl
rm
cosine_embedding_criterion_op.h
1127
0644
edit
dl
rm
cos_op.h
705
0644
edit
dl
rm
counter_ops.h
4596
0644
edit
dl
rm
create_scope_op.h
5232
0644
edit
dl
rm
cross_entropy_op.h
4420
0644
edit
dl
rm
ctc_beam_search_decoder_op.h
1102
0644
edit
dl
rm
ctc_greedy_decoder_op.h
817
0644
edit
dl
rm
cube_op.h
723
0644
edit
dl
rm
dataset_ops.h
5501
0644
edit
dl
rm
data_couple.h
464
0644
edit
dl
rm
deform_conv_op.h
3543
0644
edit
dl
rm
deform_conv_op_impl.h
13171
0644
edit
dl
rm
dense_vector_to_id_list_op.h
1797
0644
edit
dl
rm
distance_op.h
8419
0644
edit
dl
rm
do_op.h
6981
0644
edit
dl
rm
dropout_op.h
1516
0644
edit
dl
rm
elementwise_add_op.h
2024
0644
edit
dl
rm
elementwise_div_op.h
1224
0644
edit
dl
rm
elementwise_linear_op.h
1170
0644
edit
dl
rm
elementwise_logical_ops.h
5083
0644
edit
dl
rm
elementwise_mul_op.h
1224
0644
edit
dl
rm
elementwise_ops.h
19115
0644
edit
dl
rm
elementwise_ops_utils.h
1008
0644
edit
dl
rm
elementwise_op_test.h
9237
0644
edit
dl
rm
elementwise_sub_op.h
2025
0644
edit
dl
rm
elu_op.h
875
0644
edit
dl
rm
enforce_finite_op.h
2303
0644
edit
dl
rm
ensure_clipped_op.h
1608
0644
edit
dl
rm
ensure_cpu_output_op.h
1465
0644
edit
dl
rm
erf_op.h
751
0644
edit
dl
rm
expand_op.h
3877
0644
edit
dl
rm
expand_squeeze_dims_op.h
3451
0644
edit
dl
rm
exp_op.h
425
0644
edit
dl
rm
fc_inference.h
775
0644
edit
dl
rm
feature_maps_ops.h
32437
0644
edit
dl
rm
feed_blob_op.h
802
0644
edit
dl
rm
filler_op.h
18431
0644
edit
dl
rm
find_duplicate_elements_op.h
1563
0644
edit
dl
rm
find_op.h
2055
0644
edit
dl
rm
flatten_op.h
1525
0644
edit
dl
rm
flexible_top_k.h
936
0644
edit
dl
rm
floor_op.h
788
0644
edit
dl
rm
free_op.h
777
0644
edit
dl
rm
fully_connected_op.h
9351
0644
edit
dl
rm
fused_rowwise_8bit_conversion_ops.h
6601
0644
edit
dl
rm
fused_rowwise_nbitfake_conversion_ops.h
4375
0644
edit
dl
rm
fused_rowwise_nbit_conversion_ops.h
8723
0644
edit
dl
rm
fused_rowwise_random_quantization_ops.h
2607
0644
edit
dl
rm
gather_fused_8bit_rowwise_op.h
2179
0644
edit
dl
rm
gather_op.h
7505
0644
edit
dl
rm
gather_ranges_to_dense_op.h
8188
0644
edit
dl
rm
gelu_op.h
1452
0644
edit
dl
rm
generate_proposals_op.h
6256
0644
edit
dl
rm
generate_proposals_op_util_boxes.h
14309
0644
edit
dl
rm
generate_proposals_op_util_nms.h
26214
0644
edit
dl
rm
generate_proposals_op_util_nms_gpu.h
2128
0644
edit
dl
rm
given_tensor_byte_string_to_uint8_fill_op.h
2150
0644
edit
dl
rm
given_tensor_fill_op.h
3002
0644
edit
dl
rm
glu_op.h
1458
0644
edit
dl
rm
group_norm_op.h
8967
0644
edit
dl
rm
gru_unit_op.h
6626
0644
edit
dl
rm
half_float_ops.h
2732
0644
edit
dl
rm
hard_sigmoid_op.h
994
0644
edit
dl
rm
heatmap_max_keypoint_op.h
939
0644
edit
dl
rm
histogram_op.h
2421
0644
edit
dl
rm
h_softmax_op.h
4954
0644
edit
dl
rm
if_op.h
1764
0644
edit
dl
rm
im2col_op.h
8943
0644
edit
dl
rm
index_hash_ops.h
2232
0644
edit
dl
rm
index_ops.h
3155
0644
edit
dl
rm
inference_lstm_op.h
9881
0644
edit
dl
rm
instance_norm_op.h
7441
0644
edit
dl
rm
integral_image_op.h
923
0644
edit
dl
rm
is_empty_op.h
558
0644
edit
dl
rm
jsd_op.h
721
0644
edit
dl
rm
key_split_ops.h
1400
0644
edit
dl
rm
layer_norm_op.h
8098
0644
edit
dl
rm
leaky_relu_op.h
1111
0644
edit
dl
rm
lengths_pad_op.h
2574
0644
edit
dl
rm
lengths_reducer_fused_8bit_rowwise_ops.h
5532
0644
edit
dl
rm
lengths_reducer_fused_nbit_rowwise_ops.h
23465
0644
edit
dl
rm
lengths_reducer_ops.h
23315
0644
edit
dl
rm
lengths_reducer_rowwise_8bit_ops.h
6180
0644
edit
dl
rm
lengths_tile_op.h
582
0644
edit
dl
rm
lengths_top_k_op.h
1358
0644
edit
dl
rm
length_split_op.h
2259
0644
edit
dl
rm
listwise_l2r_op.h
1677
0644
edit
dl
rm
load_save_op.h
14091
0644
edit
dl
rm
load_save_op_util.h
1642
0644
edit
dl
rm
locally_connected_op.h
3872
0644
edit
dl
rm
locally_connected_op_impl.h
26495
0644
edit
dl
rm
locally_connected_op_util.h
1332
0644
edit
dl
rm
local_response_normalization_op.h
2804
0644
edit
dl
rm
log1p_op.h
717
0644
edit
dl
rm
logit_op.h
1129
0644
edit
dl
rm
log_op.h
431
0644
edit
dl
rm
loss_op.h
1058
0644
edit
dl
rm
lpnorm_op.h
1279
0644
edit
dl
rm
lstm_unit_op.h
6733
0644
edit
dl
rm
lstm_utils.h
9424
0644
edit
dl
rm
map_ops.h
8011
0644
edit
dl
rm
margin_ranking_criterion_op.h
1113
0644
edit
dl
rm
matmul_op.h
2843
0644
edit
dl
rm
max_pool_with_index_gpu.h
1155
0644
edit
dl
rm
mean_op.h
3252
0644
edit
dl
rm
merge_id_lists_op.h
2570
0644
edit
dl
rm
minmax_ops.h
3829
0644
edit
dl
rm
mish_op.h
794
0644
edit
dl
rm
mod_op.h
984
0644
edit
dl
rm
moments_op.h
4051
0644
edit
dl
rm
multi_class_accuracy_op.h
539
0644
edit
dl
rm
negate_gradient_op.h
566
0644
edit
dl
rm
negative_op.h
451
0644
edit
dl
rm
ngram_ops.h
2644
0644
edit
dl
rm
normalize_l1_op.h
1075
0644
edit
dl
rm
normalize_op.h
3013
0644
edit
dl
rm
no_default_engine_op.h
1063
0644
edit
dl
rm
numpy_tile_op.h
3643
0644
edit
dl
rm
one_hot_ops.h
2562
0644
edit
dl
rm
onnx_while_op.h
10655
0644
edit
dl
rm
operator_fallback_gpu.h
4155
0644
edit
dl
rm
op_utils_cudnn.h
2112
0644
edit
dl
rm
order_switch_ops.h
2149
0644
edit
dl
rm
pack_rnn_sequence_op.h
3074
0644
edit
dl
rm
pack_segments.h
2729
0644
edit
dl
rm
pad_op.h
2902
0644
edit
dl
rm
partition_ops.h
9958
0644
edit
dl
rm
percentile_op.h
1009
0644
edit
dl
rm
perplexity_op.h
447
0644
edit
dl
rm
piecewise_linear_transform_op.h
8281
0644
edit
dl
rm
pool_op.h
8525
0644
edit
dl
rm
pool_op_util.h
1105
0644
edit
dl
rm
pow_op.h
4677
0644
edit
dl
rm
prefetch_op.h
4661
0644
edit
dl
rm
prelu_op.h
1067
0644
edit
dl
rm
prepend_dim_op.h
2760
0644
edit
dl
rm
quantile_op.h
4120
0644
edit
dl
rm
quant_decode_op.h
5370
0644
edit
dl
rm
rank_loss_op.h
820
0644
edit
dl
rm
reciprocal_op.h
721
0644
edit
dl
rm
reducer_functors.h
24556
0644
edit
dl
rm
reduce_front_back_max_ops.h
4399
0644
edit
dl
rm
reduce_front_back_sum_mean_ops.h
5337
0644
edit
dl
rm
reduce_ops.h
9962
0644
edit
dl
rm
reduction_ops.h
5944
0644
edit
dl
rm
relu_n_op.h
990
0644
edit
dl
rm
relu_op.h
624
0644
edit
dl
rm
remove_data_blocks_op.h
2651
0644
edit
dl
rm
replace_nan_op.h
1170
0644
edit
dl
rm
reshape_op.h
5723
0644
edit
dl
rm
resize_3d_op.h
2677
0644
edit
dl
rm
resize_op.h
2307
0644
edit
dl
rm
reverse_packed_segs_op.h
2772
0644
edit
dl
rm
rmac_regions_op.h
708
0644
edit
dl
rm
rms_norm_op.h
2968
0644
edit
dl
rm
roi_align_gradient_op.h
1486
0644
edit
dl
rm
roi_align_op.h
2857
0644
edit
dl
rm
roi_align_rotated_gradient_op.h
1369
0644
edit
dl
rm
roi_align_rotated_op.h
1636
0644
edit
dl
rm
roi_pool_op.h
2503
0644
edit
dl
rm
rowmul_op.h
1947
0644
edit
dl
rm
rsqrt_op.h
729
0644
edit
dl
rm
scale_blobs_op.h
1458
0644
edit
dl
rm
scale_op.h
1019
0644
edit
dl
rm
segment_reduction_op.h
71022
0644
edit
dl
rm
self_binning_histogram_op.h
6258
0644
edit
dl
rm
selu_op.h
1545
0644
edit
dl
rm
sequence_ops.h
8264
0644
edit
dl
rm
shape_op.h
1638
0644
edit
dl
rm
sigmoid_op.h
639
0644
edit
dl
rm
sinh_op.h
711
0644
edit
dl
rm
sinusoid_position_encoding_op.h
2834
0644
edit
dl
rm
sin_op.h
705
0644
edit
dl
rm
slice_op.h
10071
0644
edit
dl
rm
softmax_op.h
1174
0644
edit
dl
rm
softmax_utils.h
447
0644
edit
dl
rm
softmax_with_loss_op.h
2883
0644
edit
dl
rm
softplus_op.h
781
0644
edit
dl
rm
softsign_op.h
675
0644
edit
dl
rm
space_batch_op.h
6848
0644
edit
dl
rm
sparse_dropout_with_replacement_op.h
1122
0644
edit
dl
rm
sparse_itemwise_dropout_with_replacement_op.h
1163
0644
edit
dl
rm
sparse_lp_regularizer_op.h
1130
0644
edit
dl
rm
sparse_normalize_op.h
834
0644
edit
dl
rm
sparse_to_dense_mask_op.h
10051
0644
edit
dl
rm
sparse_to_dense_op.h
3977
0644
edit
dl
rm
spatial_batch_norm_op.h
15175
0644
edit
dl
rm
spatial_softmax_with_loss_op.h
2182
0644
edit
dl
rm
sqrt_op.h
448
0644
edit
dl
rm
sqr_op.h
431
0644
edit
dl
rm
square_root_divide_op.h
1857
0644
edit
dl
rm
stats_put_ops.h
2813
0644
edit
dl
rm
stop_gradient.h
548
0644
edit
dl
rm
string_ops.h
2067
0644
edit
dl
rm
stump_func_op.h
2112
0644
edit
dl
rm
summarize_op.h
1875
0644
edit
dl
rm
swish_op.h
772
0644
edit
dl
rm
tanh_op.h
723
0644
edit
dl
rm
tan_op.h
705
0644
edit
dl
rm
tensor_protos_db_input.h
3633
0644
edit
dl
rm
text_file_reader_utils.h
2900
0644
edit
dl
rm
thresholded_relu_op.h
1137
0644
edit
dl
rm
tile_op.h
8741
0644
edit
dl
rm
top_k.h
1061
0644
edit
dl
rm
transpose_op.h
2082
0644
edit
dl
rm
tt_linear_op.h
6501
0644
edit
dl
rm
unique_ops.h
1666
0644
edit
dl
rm
unsafe_coalesce.h
2481
0644
edit
dl
rm
upsample_op.h
2246
0644
edit
dl
rm
utility_ops.h
49994
0644
edit
dl
rm
variable_length_sequence_padding.h
1378
0644
edit
dl
rm
weighted_multi_sampling_op.h
602
0644
edit
dl
rm
weighted_sample_op.h
739
0644
edit
dl
rm
while_op.h
1961
0644
edit
dl
rm
zero_gradient_op.h
347
0644
edit
dl
rm
Edit:
/usr/local/lib64/python3.6/site-packages/torch/include/caffe2/operators/lengths_reducer_ops.h
(23315B)
#pragma once #include "caffe2/core/context.h" #include "caffe2/core/operator.h" #include "caffe2/perfkernels/embedding_lookup.h" #ifdef USE_FBGEMM #include "fbgemm/Fbgemm.h" #endif #include <algorithm> #include <functional> namespace caffe2 { // A templated class that implements SparseLengths[Sum,WeightedSum,Mean]. template < typename T, // output type class InputTypes, // supported input types, such as TensorTypes<float> bool USE_WEIGHT = false, // Whether it is SparseLengthsWeightedSum bool USE_MEAN = false, // Whether this is SparseLengthsMean bool USE_POSITIONAL_WEIGHT = false // USE_WEIGHT = true and USE_POSITIONAL_WEIGHT = true // -> SparseLengthsPositionalWeightedSum > class CPUSparseLengthsReductionOp : public Operator<CPUContext> { public: USE_OPERATOR_FUNCTIONS(CPUContext); template <class... Args> explicit CPUSparseLengthsReductionOp(Args&&... args) : Operator<CPUContext>(std::forward<Args>(args)...) { static_assert( !(USE_WEIGHT & USE_MEAN), "Cannot both specify weight and mean."); } ~CPUSparseLengthsReductionOp() {} // Currently, we support float and at::Half inputs for input data type, and // int32_t and int64_t for the index type. bool RunOnDevice() override { return DispatchHelper<InputTypes>::call(this, Input(DATA)); } template <typename InputType> bool DoRunWithType() { return DispatchHelper<TensorTypes2<int32_t, int64_t>, InputType>::call( this, Input(INDICES)); } template <typename InputType, typename IndexType> bool DoRunWithType2() { auto& dataInput = Input(DATA); auto& indicesInput = Input(INDICES); auto& lengthsInput = Input(LENGTHS); const int64_t M = lengthsInput.size(0); const int64_t indices_size = indicesInput.numel(); auto shape = dataInput.sizes().vec(); shape[0] = M; auto* output = Output(0, shape, at::dtype<T>()); T* out_data = output->template mutable_data<T>(); if (indices_size == 0) { if (M > 0) { memset(out_data, 0, output->numel() * sizeof(T)); } return true; } CAFFE_ENFORCE_EQ(1, indicesInput.dim(), "INDICES must be a vector"); CAFFE_ENFORCE_EQ(1, lengthsInput.dim(), "LENGTHS must be a vector"); const int64_t N = dataInput.size(0); const int D = dataInput.size_from_dim(1); const InputType* in_data = dataInput.template data<InputType>(); const IndexType* indices = indicesInput.template data<IndexType>(); const int* lengths = lengthsInput.template data<int>(); const T* in_weight = nullptr; if (USE_WEIGHT) { // static if auto& weightInput = Input(WEIGHT); CAFFE_ENFORCE_EQ(1, weightInput.dim(), "WEIGHT must be a vector"); if (!USE_POSITIONAL_WEIGHT) { CAFFE_ENFORCE_EQ( weightInput.numel(), indices_size, "Weight should have the same length as indices."); } in_weight = weightInput.template data<T>(); } #ifdef USE_FBGEMM // If this is the first call or block size has changed (should never // happen actually), generate a kernel. if (D != last_block_size) { last_block_size = D; if (std::is_same<InputType, float>::value) { if (std::is_same<IndexType, std::int32_t>::value) { kernel_fp32_i32_ = fbgemm::GenerateEmbeddingSpMDM<float, std::int32_t>( D, USE_WEIGHT, USE_MEAN, /*prefetch distance*/ 16, USE_POSITIONAL_WEIGHT, /*use_offsets*/ false); } else { CAFFE_ENFORCE((std::is_same<IndexType, std::int64_t>::value)); kernel_fp32_i64_ = fbgemm::GenerateEmbeddingSpMDM<float, std::int64_t>( D, USE_WEIGHT, USE_MEAN, /*prefetch distance*/ 16, USE_POSITIONAL_WEIGHT, /*use_offsets*/ false); } } else { CAFFE_ENFORCE((std::is_same<InputType, at::Half>::value)); if (std::is_same<IndexType, std::int32_t>::value) { kernel_fp16_i32_ = fbgemm::GenerateEmbeddingSpMDM<fbgemm::float16, std::int32_t>( D, USE_WEIGHT, USE_MEAN, /*prefetch distance*/ 16, USE_POSITIONAL_WEIGHT, /*use_offsets*/ false); } else { CAFFE_ENFORCE((std::is_same<IndexType, std::int64_t>::value)); kernel_fp16_i64_ = fbgemm::GenerateEmbeddingSpMDM<fbgemm::float16, std::int64_t>( D, USE_WEIGHT, USE_MEAN, /*prefetch distance*/ 16, USE_POSITIONAL_WEIGHT, /*use_offsets*/ false); } } } bool success; if (std::is_same<InputType, float>::value) { if (std::is_same<IndexType, std::int32_t>::value) { success = kernel_fp32_i32_( M, indices_size, N, reinterpret_cast<const float*>(in_data), indicesInput.template data<std::int32_t>(), lengths, in_weight, out_data); } else { success = kernel_fp32_i64_( M, indices_size, N, reinterpret_cast<const float*>(in_data), indicesInput.template data<std::int64_t>(), lengths, in_weight, out_data); } } else { if (std::is_same<IndexType, std::int32_t>::value) { success = kernel_fp16_i32_( M, indices_size, N, reinterpret_cast<const fbgemm::float16*>(in_data), indicesInput.template data<std::int32_t>(), lengths, in_weight, out_data); } else { success = kernel_fp16_i64_( M, indices_size, N, reinterpret_cast<const fbgemm::float16*>(in_data), indicesInput.template data<std::int64_t>(), lengths, in_weight, out_data); } } if (success) { return true; } int64_t current = 0; for (int m = 0; m < M; ++m) { for (int i = 0; i < lengths[m]; ++i) { CAFFE_ENFORCE_LT( current, indices_size, "Your input seems to be incorrect: the sum of lengths values " "should be the size of the indices tensor, but it appears not."); IndexType idx = indices[current]; CAFFE_ENFORCE( 0 <= idx && idx < N, "Index ", current, " is out of bounds: ", idx, ", range 0 to ", N, ", actual batch length is ", M); ++current; } } CAFFE_ENFORCE_EQ( current, indices_size, "Your input seems to be incorrect: the sum of lengths values should be " "the size of the indices tensor, but it appears not."); return false; #endif // delegate work to perfkernel that branches based on architecture EmbeddingLookup<IndexType, InputType, T, USE_POSITIONAL_WEIGHT>( D, M, indices_size, N, in_data, indices, lengths, in_weight, nullptr, // scale_bias field is only used in SparseLengths8BitsRowwiseOp USE_MEAN, out_data); return true; } enum { DATA = 0, // Data input. WEIGHT = 1, // Weight input used in SparseLengthsWeightedSum INDICES = 1 + USE_WEIGHT, // 1 in SparseLengths[Sum,Mean] and // 2 in SparseLengthsWeightedSum LENGTHS = 2 + USE_WEIGHT, // 2 in SparseLengths[Sum, Mean], // 3 in SparseLengthsWeightedSum }; #ifdef USE_FBGEMM private: std::int64_t last_block_size{-1}; fbgemm::EmbeddingSpMDMKernelSignature<float, std::int32_t>::Type kernel_fp32_i32_; fbgemm::EmbeddingSpMDMKernelSignature<float, std::int64_t>::Type kernel_fp32_i64_; fbgemm::EmbeddingSpMDMKernelSignature<fbgemm::float16, std::int32_t>::Type kernel_fp16_i32_; fbgemm::EmbeddingSpMDMKernelSignature<fbgemm::float16, std::int64_t>::Type kernel_fp16_i64_; #endif }; template <typename T, class Context, class Engine = DefaultEngine> class TTSparseLengthsSumOp final : public Operator<Context> { public: USE_OPERATOR_CONTEXT_FUNCTIONS; template <class... Args> explicit TTSparseLengthsSumOp(Args&&... args) : Operator<Context>(std::forward<Args>(args)...), factor_i(this->template GetRepeatedArgument<int>( "factor_i", vector<int>{1, 1, 1})), factor_j(this->template GetRepeatedArgument<int>( "factor_j", vector<int>{1, 1, 1})), ranks(this->template GetRepeatedArgument<int>( "ranks", vector<int>{1, 1, 1, 1})), emb_size(this->template GetSingleArgument<int>("emb_size", 64)) { // cumprod of i, used for index slice l_cumprod.push_back(1); for (size_t i = 1; i < factor_i.size(); ++i) { l_cumprod.push_back(l_cumprod[i - 1] * factor_i[i - 1]); } } ~TTSparseLengthsSumOp() {} void Ind2Sub(int64_t* out_factor_index, const int64_t* indices, int len) { // TODO: vectorization auto N = factor_i.size(); for (int j = 0; j < len; j++) { auto idx = indices[j]; for (int i = N; i > 0; i--) { out_factor_index[j * N + i - 1] = idx / l_cumprod[i - 1]; idx = idx % l_cumprod[i - 1]; } } } bool GetSlice( std::vector<std::vector<T>>& tgt_slice, const T* core, const vector<int64_t>& ind_slice, int bs, int idx) { // implement the functinality index_select(core, 1, ind_slice) auto num_of_elements = ranks[idx] * factor_j[idx] * ranks[idx + 1]; for (int i = 0; i < bs; i++) { memcpy( tgt_slice[i].data(), core + ind_slice[i] * num_of_elements, num_of_elements * sizeof(T)); } return true; } // ind: it stores the index to each tensor core // bs: the number of indices // GatherAllRows uses two steps to calculate the lengthsum functionality: 1) it uses tensor train // to calculate the embedding for each index. 2) it sums the embedding for each bag. // In Step 1), it batches all the indices together. Specifically, for every index, it uses the pre-computed // ind of each tensor core to extract the corresponding slice of the core. Then it does gemm operation // sequentially on the slices to produce the embedding result for each index. // In Step 2), it takes the embedding computed in step 1) and apply the sum operation for each bag. bool GatherAllRows( int64_t* ind, int bs, int x_len, vector<const T*> cores, int segments, const int* lengths, T* out_data) { // compute the largest memory consumption of intermediate result // TODO: dynamic allocation size: cur_rows*factor_j[i]*ranks[i+1] // and also explore the contiguous memory storage for res and int_res int max_rank = *max_element(ranks.begin(), ranks.end()); std::vector<std::vector<T>> res(bs, std::vector<T>(emb_size * max_rank, 0)); std::vector<std::vector<T>> int_res( bs, std::vector<T>(emb_size * max_rank, 0)); // Store the matrix A vector<T*> Y_ptr(bs); // Store the intermediate result in each layer vector<T*> Z_ptr(bs); for (int b = 0; b < bs; b++) { Y_ptr[b] = res[b].data(); Z_ptr[b] = int_res[b].data(); } vector<int64_t> ind_slice(bs); int rows = 0; for (int i = 0; i < x_len; i++) { // slice cur for (int j = 0; j < bs; j++) { ind_slice[j] = ind[x_len * j + i]; } if (i == 0) { GetSlice(res, cores[i], ind_slice, bs, i); rows = factor_j[0]; } else { std::vector<std::vector<T>> slice( bs, std::vector<T>(ranks[i] * factor_j[i] * ranks[i + 1], 0)); vector<const T*> X_ptr(bs); for (int b = 0; b < bs; b++) { X_ptr[b] = slice[b].data(); } GetSlice(slice, cores[i], ind_slice, bs, i); math::GemmBatched<T, CPUContext>( CblasNoTrans, CblasNoTrans, bs, rows, factor_j[i] * ranks[i + 1], ranks[i], 1.0f, const_cast<const T**>(Y_ptr.data()), X_ptr.data(), 0.0f, Z_ptr.data(), &context_); for (int b = 0; b < bs; b++) { std::memcpy(Y_ptr[b], Z_ptr[b], (emb_size * max_rank) * sizeof(T)); } rows *= factor_j[i]; } // save the intermediate output for backward path // shape for the core auto shape = vector<int64_t>({bs, rows, ranks[i + 1]}); if (i < 2) { auto* core_data = Output(i + 1, shape, at::dtype<T>()); T* out_core = core_data->template mutable_data<T>(); for (int b = 0; b < bs; b++) { std::memcpy( out_core + b * rows * ranks[i + 1], Y_ptr[b], rows * ranks[i + 1] * sizeof(T)); } } } // reduction and store back to output vector<int64_t> cum_lengths(segments); for (int seg = 0; seg < segments; seg++) { cum_lengths[seg] = seg == 0 ? lengths[0] : lengths[seg] + cum_lengths[seg - 1]; } int length_idx = 0; vector<T> tmp_sum(emb_size, 0.0f); for (int i = 0; i <= bs; i++) { while ((length_idx < segments) && (i == cum_lengths[length_idx])) { // store the tmp_sum into output memcpy( &out_data[length_idx * emb_size], tmp_sum.data(), emb_size * sizeof(T)); length_idx++; fill(tmp_sum.begin(), tmp_sum.end(), 0.0f); } if (i == bs) { break; } transform( res[i].begin(), res[i].begin() + emb_size, tmp_sum.begin(), tmp_sum.begin(), std::plus<T>()); } return true; } bool RunOnDevice() override { const auto& dataInput0 = Input(0); const auto& dataInput1 = Input(1); const auto& dataInput2 = Input(2); const auto& indicesInput = Input(3); const auto& lengthsInput = Input(4); CAFFE_ENFORCE_EQ(1, indicesInput.dim(), "INDICES must be a vector"); CAFFE_ENFORCE_EQ(1, lengthsInput.dim(), "LENGTHS must be a vector"); int N = factor_i.size(); const int64_t M = lengthsInput.size(0); auto shape = vector<int64_t>({M, emb_size}); auto* output = Output(0, shape, at::dtype<T>()); T* out_data = output->template mutable_data<T>(); const T* core0 = dataInput0.template data<T>(); const T* core1 = dataInput1.template data<T>(); const T* core2 = dataInput2.template data<T>(); const int* lengths = lengthsInput.template data<int>(); vector<const T*> cores = {core0, core1, core2}; const int64_t* indices = indicesInput.template data<int64_t>(); // Store the factor index for backward path auto index_shape = vector<int64_t>({indicesInput.size(), N}); auto* index_data = Output(3, index_shape, at::dtype<int64_t>()); int64_t* out_factor_index = index_data->template mutable_data<int64_t>(); // Store the factorized index for each core Ind2Sub(out_factor_index, indices, indicesInput.size()); return GatherAllRows( out_factor_index, indicesInput.size(), N, cores, M, lengths, out_data); } protected: vector<int> factor_i; vector<int> factor_j; vector<int> ranks; vector<int> l_cumprod; int emb_size; }; template <typename T, class Context> class TTSparseLengthsSumGradientOp final : public Operator<Context> { public: USE_OPERATOR_CONTEXT_FUNCTIONS; template <class... Args> explicit TTSparseLengthsSumGradientOp(Args&&... args) : Operator<Context>(std::forward<Args>(args)...) {} bool RunOnDevice() override; ~TTSparseLengthsSumGradientOp() {} }; // implement the graident op for TTLengthSumGradient op template <typename T, class Context> bool TTSparseLengthsSumGradientOp<T, Context>::RunOnDevice() { const auto& core0 = Input(0); const auto& core1 = Input(1); const auto& core2 = Input(2); const auto& lengths = Input(3); const auto& core0_out = Input(4); const auto& core1_out = Input(5); const auto& index_out = Input(6); const auto& dY = Input(7); const int* lengths_data = lengths.template data<int>(); const T* dY_data = dY.template data<T>(); // restore the arguments from shape const int64_t bs = index_out.size(0); const int64_t emb_size = dY.size(1); const int64_t num_segments = lengths.size(0); auto core0_shape = core0.sizes().vec(); auto core1_shape = core1.sizes().vec(); auto core2_shape = core2.sizes().vec(); auto core0_out_shape = core0_out.sizes().vec(); auto core1_out_shape = core1_out.sizes().vec(); auto* dCore0 = Output(0, core0_shape, at::dtype<T>()); auto* dCore1 = Output(1, core1_shape, at::dtype<T>()); auto* dCore2 = Output(2, core2_shape, at::dtype<T>()); T* dCore0_data = dCore0->template mutable_data<T>(); T* dCore1_data = dCore1->template mutable_data<T>(); T* dCore2_data = dCore2->template mutable_data<T>(); memset( dCore0_data, 0.0f, sizeof(T) * accumulate( core0_shape.begin(), core0_shape.end(), 1, std::multiplies<T>())); memset( dCore1_data, 0.0f, sizeof(T) * accumulate( core1_shape.begin(), core1_shape.end(), 1, std::multiplies<T>())); memset( dCore2_data, 0.0f, sizeof(T) * accumulate( core2_shape.begin(), core2_shape.end(), 1, std::multiplies<T>())); int64_t* index_out_data = index_out.template mutable_data<int64_t>(); vector<vector<int64_t>> index_slice(bs, vector<int64_t>(3, 0)); for (int64_t b = 0; b < bs; b++) { memcpy(index_slice[b].data(), index_out_data + b * 3, 3 * sizeof(int64_t)); } vector<const T*> A_ptr(bs); vector<T*> B_ptr(bs); vector<T*> C_ptr(bs); // size of each batch int64_t num_of_elements = 0; // construct the ranks // expand the gradient into all indices vector<vector<T>> core2_out_grad(bs, vector<T>(emb_size, 0)); int64_t data_index = 0; for (int64_t range_index = 0; range_index < num_segments; ++range_index) { for (int64_t start = data_index; data_index < start + lengths_data[range_index]; ++data_index) { memcpy( core2_out_grad[data_index].data(), dY_data + range_index * emb_size, emb_size * sizeof(T)); } } // ======================================================= // Calculate dCore2_data: // 1) Transpose core1_out and multiply iwth core2_out_grad // 2) add to dCore2_data vector<vector<T>> dCore2_data_slice_grad( bs, vector<T>(core2_shape[1] * core2_shape[2] * core2_shape[3], 0)); const T* core1_out_data = core1_out.template data<T>(); // const T* core1_out_p[bs]; for (int64_t b = 0; b < bs; b++) { A_ptr[b] = core1_out_data + b * core1_out.size(1) * core1_out.size(2); B_ptr[b] = core2_out_grad[b].data(); C_ptr[b] = dCore2_data_slice_grad[b].data(); } math::GemmBatched<T, CPUContext>( CblasTrans, CblasNoTrans, bs, core2.size(1), // M core2.size(2) * core2.size(3), // N core1_out.size(1), // K 1.0f, const_cast<const T**>(A_ptr.data()), const_cast<const T**>(B_ptr.data()), 0.0f, C_ptr.data(), &context_); // update the corresponding slice num_of_elements = core2_shape[1] * core2_shape[2] * core2_shape[3]; T* core2_data = core2.template mutable_data<T>(); vector<vector<T>> core2_slice( bs, vector<T>(core2_shape[1] * core2_shape[2] * core2_shape[3], 0)); for (int64_t b = 0; b < bs; b++) { for (int i = 0; i < num_of_elements; i++) { dCore2_data[index_slice[b][2] * num_of_elements + i] += C_ptr[b][i]; } memcpy( core2_slice[b].data(), core2_data + index_slice[b][2] * num_of_elements, sizeof(T) * num_of_elements); } // Calculate core1_out_grad vector<vector<T>> core1_out_grad( bs, vector<T>(core1_out_shape[1] * core1_out_shape[2], 0)); for (int64_t b = 0; b < bs; b++) { A_ptr[b] = core2_out_grad[b].data(); B_ptr[b] = core2_slice[b].data(); C_ptr[b] = core1_out_grad[b].data(); } math::GemmBatched<T, CPUContext>( CblasNoTrans, CblasTrans, bs, core1_out.size(1), // M core2_shape[1], // N core2_shape[2] * core2_shape[3], // K 1.0f, const_cast<const T**>(A_ptr.data()), const_cast<const T**>(B_ptr.data()), 0.0f, C_ptr.data(), &context_); // ======================================================= // Calcuate dCore1_data: // 1) Transpose core1_out_grad and multiply with core0_out // 2) Transpose the result and then add to dCore1_data vector<vector<T>> dCore1_data_slice_grad( bs, vector<T>(core1_shape[1] * core1_shape[2] * core1_shape[3], 0)); const T* core0_out_data = core0_out.template data<T>(); for (int64_t b = 0; b < bs; b++) { A_ptr[b] = core0_out_data + b * core0_out.size(1) * core0_out.size(2); B_ptr[b] = core1_out_grad[b].data(); C_ptr[b] = dCore1_data_slice_grad[b].data(); } math::GemmBatched<T, CPUContext>( CblasTrans, CblasNoTrans, bs, core1.size(1), // M core1.size(2) * core1.size(3), // N core0_out.size(1), // K 1.0f, const_cast<const T**>(A_ptr.data()), const_cast<const T**>(B_ptr.data()), 0.0f, C_ptr.data(), &context_); // update the corresponding slice num_of_elements = core1_shape[1] * core1_shape[2] * core1_shape[3]; T* core1_data = core1.template mutable_data<T>(); vector<vector<T>> core1_slice( bs, vector<T>(core1_shape[1] * core1_shape[2] * core1_shape[3], 0)); for (int64_t b = 0; b < bs; b++) { for (int i = 0; i < num_of_elements; i++) { dCore1_data[index_slice[b][1] * num_of_elements + i] += C_ptr[b][i]; } memcpy( core1_slice[b].data(), core1_data + index_slice[b][1] * num_of_elements, sizeof(T) * num_of_elements); } // Calcuate core0_out_grad vector<vector<T>> core0_out_grad( bs, vector<T>(core0_out_shape[1] * core0_out_shape[2], 0)); for (int64_t b = 0; b < bs; b++) { A_ptr[b] = core1_out_grad[b].data(); B_ptr[b] = core1_slice[b].data(); C_ptr[b] = core0_out_grad[b].data(); } math::GemmBatched<T, CPUContext>( CblasNoTrans, CblasTrans, bs, core0_out.size(1), // M core1_shape[1], // N core1_shape[2] * core1_shape[3], // K 1.0f, const_cast<const T**>(A_ptr.data()), const_cast<const T**>(B_ptr.data()), 0.0f, C_ptr.data(), &context_); num_of_elements = core0_shape[1] * core0_shape[2] * core0_shape[3]; for (int64_t b = 0; b < bs; b++) { for (int i = 0; i < num_of_elements; i++) { dCore0_data[index_slice[b][0] * num_of_elements + i] += C_ptr[b][i]; } } return true; } } // namespace caffe2
Save
cmd:
run