(
BART_decoder, input_ids, encoder_last_hidden_state, timing_profile, use_cuda=True, use_cache=False, past_key_values=None
)
| 48 | |
| 49 | @use_cuda |
| 50 | def 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 |
no test coverage detected