| 1506 | void SubRow(const Tensor &v, Tensor *M) { AddRow(-1, 1, v, M); } |
| 1507 | |
| 1508 | void SumColumns(const Tensor &M, Tensor *v) { |
| 1509 | if (M.transpose()) { |
| 1510 | Tensor X = Transpose(M); |
| 1511 | SumRows(X, v); |
| 1512 | } else { |
| 1513 | CHECK_EQ(M.nDim(), 2u); |
| 1514 | // CHECK_EQ(v->nDim(), 1u); (chonho) shape of v is 2-element tuple |
| 1515 | size_t nb_row = M.shape().at(0), nb_col = M.shape().at(1); |
| 1516 | CHECK_EQ(nb_row, v->Size()); |
| 1517 | |
| 1518 | Tensor one(Shape{nb_col}, M.device(), M.data_type()); |
| 1519 | one.SetValue(1.0f); // TODO(wangwei) cast type |
| 1520 | Mult(M, one, v); |
| 1521 | } |
| 1522 | } |
| 1523 | void SumRows(const Tensor &M, Tensor *v) { |
| 1524 | if (M.transpose()) { |
| 1525 | Tensor X = Transpose(M); |