MCPcopy Create free account
hub / github.com/DeepRec-AI/DeepRec / XlaTensorFormat

Function XlaTensorFormat

tensorflow/compiler/tf2xla/kernels/pooling_ops.cc:136–149  ·  view source on GitHub ↗

Converts the tensor data format to the one required by the XLA pooling library.

Source from the content-addressed store, hash-verified

134// Converts the tensor data format to the one required by the XLA pooling
135// library.
136xla::TensorFormat XlaTensorFormat(tensorflow::TensorFormat data_format,
137 int num_spatial_dims) {
138 int num_dims = num_spatial_dims + 2;
139 int batch_dimension = GetTensorBatchDimIndex(num_dims, data_format);
140 int feature_dimension = GetTensorFeatureDimIndex(num_dims, data_format);
141 absl::InlinedVector<int64, 4> spatial_dimensions(num_spatial_dims);
142 for (int spatial_dim = 0; spatial_dim < num_spatial_dims; ++spatial_dim) {
143 spatial_dimensions[spatial_dim] =
144 GetTensorSpatialDimIndex(num_dims, data_format, spatial_dim);
145 }
146 return xla::TensorFormat(/*batch_dimension=*/batch_dimension,
147 /*feature_dimension=*/feature_dimension,
148 /*spatial_dimensions=*/spatial_dimensions);
149}
150
151class MaxPoolOp : public PoolingOp {
152 public:

Callers 4

CompileMethod · 0.85
CompileMethod · 0.85
CompileMethod · 0.85
CompileMethod · 0.85

Calls 4

GetTensorBatchDimIndexFunction · 0.85
GetTensorFeatureDimIndexFunction · 0.85
GetTensorSpatialDimIndexFunction · 0.85
TensorFormatClass · 0.50

Tested by

no test coverage detected