| 208 | } |
| 209 | |
| 210 | XlaOp 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 | |
| 228 | XlaOp Triangle(XlaOp x, bool lower) { |
| 229 | return lower ? Select(TriangleMask(x, 0), x, ZerosLike(x)) |
no test coverage detected