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

Function TriangleMask

tensorflow/compiler/xla/client/lib/matrix.cc:210–226  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

208}
209
210XlaOp TriangleMask(XlaOp x, int diagonal) {
211 XlaBuilder* builder = x.builder();
212 return builder->ReportErrorOrReturn([&]() -> StatusOr<XlaOp> {
213 TF_ASSIGN_OR_RETURN(Shape shape, builder->GetShape(x));
214 const int64 n_dims = shape.rank();
215 TF_RET_CHECK(n_dims >= 2);
216 const int64 m = shape.dimensions(n_dims - 2);
217 const int64 n = shape.dimensions(n_dims - 1);
218 absl::Span<const int64> major_dims =
219 AsInt64Slice(shape.dimensions()).subspan(/*pos=*/0, /*len=*/n_dims - 2);
220 auto a = Iota(builder, S32, n);
221 auto b = Iota(builder, S32, m) + ConstantR0<int32>(builder, diagonal);
222 XlaOp indicator;
223 indicator = Ge(b, Broadcast(a, {m}), /*broadcast_dimensions=*/{0});
224 return Broadcast(indicator, major_dims);
225 });
226}
227
228XlaOp Triangle(XlaOp x, bool lower) {
229 return lower ? Select(TriangleMask(x, 0), x, ZerosLike(x))

Callers 2

TriangleFunction · 0.85
BuildTriangularSolveFunction · 0.85

Calls 8

BroadcastFunction · 0.85
ReportErrorOrReturnMethod · 0.80
AsInt64SliceFunction · 0.50
IotaFunction · 0.50
GeFunction · 0.50
builderMethod · 0.45
rankMethod · 0.45
dimensionsMethod · 0.45

Tested by

no test coverage detected