(self, model, network_metadata)
| 200 | class BARTDecoderTRTEngine(TRTEngineFile): |
| 201 | |
| 202 | def __init__(self, model, network_metadata): |
| 203 | super().__init__(model, BARTDecoderConverter, network_metadata) |
| 204 | self.max_trt_workspace = BARTModelTRTConfig.MAX_DECODER_WORKSPACE_MB[network_metadata.variant] |
| 205 | |
| 206 | def get_network_definition(self, network_definition): |
| 207 | return add_extra_fp32(network_definition) |