MCPcopy Create free account
hub / github.com/NVIDIA/TensorRT / decoder_inference

Function decoder_inference

demo/HuggingFace/BART/measurements.py:50–68  ·  view source on GitHub ↗
(
    BART_decoder, input_ids, encoder_last_hidden_state, timing_profile, use_cuda=True, use_cache=False, past_key_values=None
)

Source from the content-addressed store, hash-verified

48
49@use_cuda
50def decoder_inference(
51 BART_decoder, input_ids, encoder_last_hidden_state, timing_profile, use_cuda=True, use_cache=False, past_key_values=None
52):
53 # This implementation is a bit ugly. Moving implementation of the model to check HFRunner would be cleaner.
54 if isinstance(BART_decoder, TRTNativeRunner):
55 # Function is technically in BARTTRTDecoder however due to circular import, TRTNativeRunner in this module scope
56 # implies the existence of this function.
57 BART_decoder.set_encoder_hidden_states_for_inference_cycle(encoder_last_hidden_state)
58 BART_decoder.set_return_device("cuda" if use_cuda else "cpu")
59
60 def decoder_stmt():
61 BART_decoder(
62 input_ids=input_ids, encoder_hidden_states=encoder_last_hidden_state, use_cache=use_cache,
63 past_key_values=past_key_values
64 )
65
66 decoder_e2e_time = measure_python_inference_code(decoder_stmt, timing_profile)
67
68 return (decoder_stmt(), decoder_e2e_time)
69
70
71@use_cuda

Callers 3

execute_inferenceMethod · 0.90
execute_inferenceMethod · 0.90
execute_inferenceMethod · 0.90

Tested by

no test coverage detected