MCPcopy Create free account
hub / github.com/AI-Hypercomputer/maxtext / test_basic_prefill

Method test_basic_prefill

tests/maxengine_test.py:109–135  ·  view source on GitHub ↗
(self)

Source from the content-addressed store, hash-verified

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)

Callers

nothing calls this directly

Calls 4

get_dataMethod · 0.95
prefillMethod · 0.95
MaxEngineClass · 0.90
initMethod · 0.45

Tested by

no test coverage detected