Download daemon/src/session.rs from Snapkitty/bert-agent: direct link, hf CLI and curl.
- Browser
- Download file 1.68 kB
-
https://huggingface.co/Snapkitty/bert-agent/resolve/main/daemon/src/session.rs
- Command line
-
hf download hf://Snapkitty/bert-agent/daemon/src/session.rs
-
curl -L -o session.rs https://huggingface.co/Snapkitty/bert-agent/resolve/main/daemon/src/session.rs
1.68 kB
| //! session.rs — TensorRT ORT session initialisation | |
| use crate::types::DaemonConfig; | |
| use ort::{ExecutionProvider, GraphOptimizationLevel, Session, SessionBuilder}; | |
| /// Initialise an ORT session backed by TensorrtExecutionProvider. | |
| /// | |
| /// On first call: TRT compiles the FP16 engine (~5 min). | |
| /// Subsequent calls: loads cached .plan file instantly. | |
| pub fn build_trt_session(cfg: &DaemonConfig) -> ort::Result<Session> { | |
| std::fs::create_dir_all(&cfg.trt_cache_dir) | |
| .expect("failed to create TRT cache directory"); | |
| let trt_provider = ExecutionProvider::TensorRT( | |
| ort::TensorRTExecutionProviderOptions::default() | |
| .with_device_id(0) | |
| .with_fp16_enable(true) | |
| .with_engine_cache_enable(true) | |
| .with_engine_cache_path(&cfg.trt_cache_dir) | |
| .with_profile_min_shapes(&format!( | |
| "input_ids:1x16,attention_mask:1x16,token_type_ids:1x16" | |
| )) | |
| .with_profile_opt_shapes(&format!( | |
| "input_ids:{}x128,attention_mask:{}x128,token_type_ids:{}x128", | |
| cfg.max_batch_size, cfg.max_batch_size, cfg.max_batch_size | |
| )) | |
| .with_profile_max_shapes(&format!( | |
| "input_ids:{}x512,attention_mask:{}x512,token_type_ids:{}x512", | |
| cfg.max_batch_size, cfg.max_batch_size, cfg.max_batch_size | |
| )), | |
| ); | |
| ort::init() | |
| .with_execution_providers([trt_provider]) | |
| .commit()?; | |
| SessionBuilder::new()? | |
| .with_optimization_level(GraphOptimizationLevel::Level3)? | |
| .with_intra_threads(4)? | |
| .commit_from_file(&cfg.model_path) | |
| } | |