Householder reflection on the trailing elements of a col of a matrix. After applying it on the matrix, all elements in [(i+1):, j] become zeros, i.e., H = I - beta * [1; v] * [1; v]', then, H * A[i:, j] = [xnorm, 0, 0, ..., 0]
| 172 | // H * A[i:, j] = [xnorm, 0, 0, ..., 0] |
| 173 | // |
| 174 | StatusOr<HouseHolderResult> HouseCol(XlaOp a, XlaOp i, XlaOp j, XlaOp eps, |
| 175 | PrecisionConfig::Precision precision) { |
| 176 | XlaBuilder* builder = a.builder(); |
| 177 | TF_ASSIGN_OR_RETURN(Shape a_shape, builder->GetShape(a)); |
| 178 | const int64 num_dims = a_shape.rank(); |
| 179 | const int64 m = ShapeUtil::GetDimension(a_shape, -2); |
| 180 | XlaOp zero = ScalarLike(i, 0); |
| 181 | XlaOp x = DynamicSliceInMinorDims(a, {zero, j}, {m, 1}); |
| 182 | |
| 183 | const int64 num_batch_dims = num_dims - 2; |
| 184 | std::vector<int64> batch_dims(num_batch_dims); |
| 185 | for (int k = 0; k < num_batch_dims; ++k) { |
| 186 | batch_dims[k] = ShapeUtil::GetDimension(a_shape, k); |
| 187 | } |
| 188 | |
| 189 | TF_ASSIGN_OR_RETURN(Shape x_shape, builder->GetShape(x)); |
| 190 | auto idx = Iota(builder, ShapeUtil::MakeShape(S32, x_shape.dimensions()), |
| 191 | num_dims - 2); |
| 192 | auto zeros = ZerosLike(x); |
| 193 | auto v = Select(Gt(idx, i), x, zeros); |
| 194 | |
| 195 | auto one = ScalarLike(v, 1.0); |
| 196 | |
| 197 | auto sigma = |
| 198 | Sqrt(Reduce(Square(v), ScalarLike(v, 0.0), |
| 199 | CreateScalarAddComputation(x_shape.element_type(), builder), |
| 200 | {num_dims - 2})); |
| 201 | |
| 202 | std::vector<int64> broadcast_dims(num_dims - 1); |
| 203 | std::iota(broadcast_dims.begin(), broadcast_dims.end(), 0); |
| 204 | broadcast_dims[num_dims - 2] = num_dims - 1; |
| 205 | auto x_0i = DynamicSliceInMinorDims(x, {i, zero}, {1, 1}); |
| 206 | auto mu = Mul(sigma, Sqrt(Square(Div(x_0i, sigma, broadcast_dims)) + one), |
| 207 | broadcast_dims); |
| 208 | |
| 209 | auto v_0i = Select( |
| 210 | Le(x_0i, ScalarLike(x_0i, 0.0)), Sub(x_0i, mu), |
| 211 | -Mul(sigma, Div(sigma, Add(x_0i, mu), broadcast_dims), broadcast_dims)); |
| 212 | |
| 213 | auto beta = Div(ScalarLike(v_0i, 2.0), |
| 214 | (Square(Div(sigma, v_0i, broadcast_dims)) + one)); |
| 215 | |
| 216 | v = Select( |
| 217 | BroadcastInDim(Lt(sigma, eps), x_shape.dimensions(), broadcast_dims), v, |
| 218 | v / v_0i); |
| 219 | v = Select(Eq(idx, i), zeros + one, v); |
| 220 | |
| 221 | beta = Select(Lt(Add(sigma, ZerosLike(beta), broadcast_dims), eps), |
| 222 | ZerosLike(beta), beta); |
| 223 | |
| 224 | HouseHolderResult result; |
| 225 | result.v = v; |
| 226 | result.beta = beta; |
| 227 | result.a = Sub( |
| 228 | a, Mul(beta, BatchDot(v, false, BatchDot(v, true, a, false, precision), |
| 229 | false, precision))); |
| 230 | |
| 231 | return result; |
nothing calls this directly
no test coverage detected