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

Method test_forward

axlearn/audio/decoder_asr_test.py:1155–1232  ·  view source on GitHub ↗

Tests that loss computation excludes empty sequence, and respects paddings.

(self, target_labels, tile_input)

Source from the content-addressed store, hash-verified

1153 )
1154 @set_threefry_partitionable(False) # TODO(Luzy): update for threefry_partitionable True
1155 def test_forward(self, target_labels, tile_input):
1156 """Tests that loss computation excludes empty sequence, and respects paddings."""
1157 am_dim, bos_id = 4, 1
1158 layer, layer_params, prng_key = self._set_up_transducer(vocab_size=20)
1159
1160 # Generate inputs.
1161 if tile_input:
1162 batch_size, src_len, max_src_len = 4, 5, 10
1163 # [batch_size, src_len, am_dim].
1164 inputs = jnp.tile(
1165 jax.random.normal(jax.random.PRNGKey(707), [1, max_src_len, am_dim]) * 1000,
1166 [batch_size, 1, 1],
1167 )
1168 paddings = jnp.tile(jnp.arange(max_src_len)[None, :] >= src_len, [batch_size, 1])
1169 # Generate different padding data.
1170 pad_inputs_data = (
1171 jax.random.normal(jax.random.PRNGKey(124), [batch_size, max_src_len, am_dim]) * 200
1172 )
1173 # Generate inputs with the same data at non-pad positions.
1174 inputs = jnp.where(paddings[:, :, None], pad_inputs_data, inputs)
1175 assert_allclose(
1176 jnp.diff(inputs[:, :src_len], axis=0), jnp.zeros([batch_size - 1, src_len, am_dim])
1177 )
1178 self.assertGreater(
1179 jnp.abs(jnp.diff(inputs[:, src_len:], axis=0)).sum(),
1180 1e3,
1181 )
1182 target_labels = jnp.tile(target_labels, [batch_size, 1])
1183 else:
1184 batch_size = target_labels.shape[0]
1185 max_src_len = 10
1186 src_len = np.array([10, 0, 7])
1187 # [batch_size, src_len, am_dim].
1188 inputs = (
1189 jax.random.normal(jax.random.PRNGKey(311), [batch_size, max_src_len, am_dim]) * 1000
1190 )
1191 paddings = jnp.arange(max_src_len)[None, :] >= src_len[:, None]
1192
1193 input_ids = jnp.concatenate(
1194 [jnp.full([batch_size, 1], bos_id), target_labels[:, :-1]], axis=1
1195 )
1196
1197 @jax.jit
1198 def jit_forward(input_batch):
1199 (loss, aux_outputs), _ = F(
1200 layer,
1201 inputs=dict(input_batch=input_batch),
1202 is_training=True,
1203 prng_key=prng_key,
1204 state=layer_params,
1205 )
1206 return loss, aux_outputs
1207
1208 # Compute test loss.
1209 loss, aux_outputs = jit_forward(
1210 dict(
1211 inputs=inputs,
1212 paddings=paddings,

Callers

nothing calls this directly

Calls 3

_set_up_transducerMethod · 0.95
assert_allcloseFunction · 0.90
assertNestedEqualMethod · 0.80

Tested by

no test coverage detected