Pads or truncates ids so as to divide max length, then group into temporary batch.
(ids: tf.Tensor)
| 120 | assert 0 <= max_padding_fraction <= 1.0 |
| 121 | |
| 122 | def batch(ids: tf.Tensor) -> tf.Tensor: |
| 123 | """Pads or truncates ids so as to divide max length, then group into temporary batch.""" |
| 124 | len_ids = tf.shape(ids)[0] |
| 125 | remainder = len_ids % tf.constant(max_len) |
| 126 | tf_pad_id = tf.constant(vocab.pad_id, dtype=tf.int32) |
| 127 | # If the remainder isn't long enough to satisfy max_padding_fraction for a given example, |
| 128 | # drop it, else pad to fill to max_len. |
| 129 | new_ids = tf.cond( |
| 130 | remainder > tf.constant(int(max_len * (1 - max_padding_fraction)), dtype=tf.int32), |
| 131 | lambda: tf.concat( # pylint: disable=unexpected-keyword-arg,no-value-for-parameter |
| 132 | (ids, tf.broadcast_to(tf_pad_id, shape=(tf.constant(max_len) - remainder,))), |
| 133 | axis=0, |
| 134 | ), |
| 135 | lambda: ids[: len_ids - remainder], |
| 136 | ) |
| 137 | return tf.reshape(new_ids, shape=(-1, max_len)) |
| 138 | |
| 139 | def process_batched(inputs: dict[str, Any]) -> dict[str, Any]: |
| 140 | """Chunks each jagged input window into a batch of equal-length training examples.""" |