Tests that loss computation excludes empty sequence, and respects paddings.
(self, target_labels, tile_input)
| 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, |
nothing calls this directly
no test coverage detected