* \brief Get a Tensor object representing a single batch. * \details If this tensor only has a single batch, then broadcast. Otherwise, check to make sure that the requested batch is smaller than the number of batches. * * TODO: This is a bit wasteful, as it re-calculates `bs.batch_size()` every time. * * \param b Batch id * \return Sub tensor at batch `b` */
| 90 | * \return Sub tensor at batch `b` |
| 91 | */ |
| 92 | Tensor batch_elem(unsigned b) const { |
| 93 | if (d.batch_elems() == 1) { |
| 94 | return *this; |
| 95 | } else { |
| 96 | if (b >= d.batch_elems()) { |
| 97 | std::stringstream ss; |
| 98 | ss << "Requested batch id " << b << " is greater than the number of batch " << d.batch_elems(); |
| 99 | throw std::runtime_error(ss.str()); |
| 100 | } |
| 101 | const unsigned bsize = d.batch_size(); |
| 102 | Dim new_d(d); new_d.bd = 1; |
| 103 | Tensor ret(new_d, v + bsize * b, device, mem_pool); |
| 104 | return ret; |
| 105 | } |
| 106 | } |
| 107 | |
| 108 | // get tensors for all batches |
| 109 | /** |
nothing calls this directly
no test coverage detected