()
| 17 | |
| 18 | |
| 19 | def test_ModelLoader(): |
| 20 | kv_cache_config = KvCacheConfig(free_gpu_memory_fraction=0.4) |
| 21 | |
| 22 | # Test with HF model |
| 23 | temp_dir = tempfile.TemporaryDirectory() |
| 24 | |
| 25 | def build_engine(): |
| 26 | args = TrtLlmArgs(model=llama_model_path, |
| 27 | kv_cache_config=kv_cache_config) |
| 28 | model_loader = ModelLoader(args) |
| 29 | engine_dir = model_loader(engine_dir=Path(temp_dir.name)) |
| 30 | assert engine_dir |
| 31 | return engine_dir |
| 32 | |
| 33 | # Test with engine |
| 34 | args = TrtLlmArgs(model=build_engine(), kv_cache_config=kv_cache_config) |
| 35 | assert args.model_format is _ModelFormatKind.TLLM_ENGINE |
| 36 | print(f'engine_dir: {args.model}') |
| 37 | model_loader = ModelLoader(args) |
| 38 | engine_dir = model_loader() |
| 39 | assert engine_dir == args.model |
| 40 | |
| 41 | |
| 42 | def test_CachedModelLoader(): |
nothing calls this directly
no test coverage detected