| 1300 | } // namespace |
| 1301 | |
| 1302 | bool TensorEquals(const Tensor& left, const Tensor& right, const EqualOptions& opts) { |
| 1303 | if (left.type_id() != right.type_id()) { |
| 1304 | return false; |
| 1305 | } else if (left.size() == 0 && right.size() == 0) { |
| 1306 | return true; |
| 1307 | } else if (left.shape() != right.shape()) { |
| 1308 | return false; |
| 1309 | } |
| 1310 | |
| 1311 | switch (left.type_id()) { |
| 1312 | // TODO: Support half-float tensors |
| 1313 | // case Type::HALF_FLOAT: |
| 1314 | case Type::FLOAT: |
| 1315 | return FloatTensorEquals<FloatType>(left, right, opts); |
| 1316 | |
| 1317 | case Type::DOUBLE: |
| 1318 | return FloatTensorEquals<DoubleType>(left, right, opts); |
| 1319 | |
| 1320 | default: |
| 1321 | return IntegerTensorEquals(left, right); |
| 1322 | } |
| 1323 | } |
| 1324 | |
| 1325 | namespace { |
| 1326 |
no test coverage detected