MCPcopy Create free account
hub / github.com/apple/axlearn / batch

Function batch

axlearn/common/input_lm.py:122–137  ·  view source on GitHub ↗

Pads or truncates ids so as to divide max length, then group into temporary batch.

(ids: tf.Tensor)

Source from the content-addressed store, hash-verified

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."""

Callers 2

process_batchedFunction · 0.70
test_split_prng_keyMethod · 0.70

Calls 1

shapeMethod · 0.45

Tested by 1

test_split_prng_keyMethod · 0.56