| 107 | jax.tree.map(np.testing.assert_array_equal, got_unstacked, input_d) |
| 108 | |
| 109 | def test_basic_prefill(self): |
| 110 | devices_array = maxtext_utils.create_device_mesh(self.cfg) |
| 111 | mesh = Mesh(devices_array, self.cfg.mesh_axes) |
| 112 | quant = quantizations.configure_quantization(self.cfg) |
| 113 | model = models.transformer_as_linen(config=self.cfg, mesh=mesh, quant=quant, model_mode=MODEL_MODE_PREFILL) |
| 114 | ids, decoder_segment_ids, decoder_positions = self.get_data() |
| 115 | |
| 116 | transformer_vars = model.init( |
| 117 | {"params": self.rng, "aqt": self.rng, "dropout": self.rng}, |
| 118 | ids, |
| 119 | decoder_positions, |
| 120 | decoder_segment_ids, |
| 121 | enable_dropout=False, |
| 122 | ) |
| 123 | input_tokens = jnp.array([1, 306, 5360, 304, 0, 0, 0, 0]) |
| 124 | true_length = 4 |
| 125 | engine = MaxEngine(self.cfg, jax.devices()) |
| 126 | prefill_result, first_token = engine.prefill( |
| 127 | params=transformer_vars, padded_tokens=input_tokens, true_length=true_length |
| 128 | ) |
| 129 | |
| 130 | self.assertEqual(prefill_result["generated_tokens"], jnp.array([0])) |
| 131 | # test default strategy is gready which choose only one next token |
| 132 | self.assertEqual(prefill_result["tokens"].size, 1) |
| 133 | self.assertNotEqual(prefill_result["tokens"], jnp.array([0])) |
| 134 | self.assertTrue(jnp.array_equal(first_token.data.size, 3)) |
| 135 | self.assertEqual(first_token.log_prob.shape, (1, 1)) |
| 136 | |
| 137 | def test_basic_decode(self): |
| 138 | devices_array = maxtext_utils.create_device_mesh(self.cfg) |