MCPcopy Create free account
hub / github.com/OpenNMT/CTranslate2 / test_transformers_encoder

Function test_transformers_encoder

python/tests/test_transformers.py:317–365  ·  view source on GitHub ↗
(clear_transformers_cache, tmp_dir, device, model_name)

Source from the content-addressed store, hash-verified

315 ],
316)
317def test_transformers_encoder(clear_transformers_cache, tmp_dir, device, model_name):
318 import torch
319 import transformers
320
321 text = ["Hello world!", "Hello, my dog is cute"]
322
323 model = transformers.AutoModel.from_pretrained(model_name)
324 tokenizer = transformers.AutoTokenizer.from_pretrained(model_name)
325
326 inputs = tokenizer(text, return_tensors="pt", padding=True)
327
328 inputs.to(device)
329 model.to(device)
330
331 with torch.no_grad():
332 outputs = model(**inputs)
333
334 mask = inputs.attention_mask.unsqueeze(-1).cpu().numpy()
335 ref_last_hidden_state = outputs.last_hidden_state.cpu().numpy()
336 ref_pooler_output = (
337 outputs.pooler_output.cpu().numpy()
338 if hasattr(outputs, "pooler_output")
339 else None
340 )
341
342 converter = ctranslate2.converters.TransformersConverter(model_name)
343 output_dir = str(tmp_dir.join("ctranslate2_model"))
344 output_dir = converter.convert(output_dir)
345
346 encoder = ctranslate2.Encoder(output_dir, device=device)
347
348 ids = [tokenizer(t).input_ids for t in text]
349 outputs = encoder.forward_batch(ids)
350
351 last_hidden_state = _to_numpy(outputs.last_hidden_state, device)
352 assert last_hidden_state.shape == ref_last_hidden_state.shape
353
354 last_hidden_state *= mask
355 ref_last_hidden_state *= mask
356 np.testing.assert_array_almost_equal(
357 last_hidden_state, ref_last_hidden_state, decimal=5
358 )
359
360 if ref_pooler_output is not None:
361 pooler_output = _to_numpy(outputs.pooler_output, device)
362 assert pooler_output.shape == ref_pooler_output.shape
363 np.testing.assert_array_almost_equal(
364 pooler_output, ref_pooler_output, decimal=5
365 )
366
367
368def _to_numpy(storage, device):

Callers

nothing calls this directly

Calls 6

_to_numpyFunction · 0.85
joinMethod · 0.80
toMethod · 0.45
numpyMethod · 0.45
convertMethod · 0.45
forward_batchMethod · 0.45

Tested by

no test coverage detected