diff --git a/tests/integration/defs/perf/pytorch_model_config.py b/tests/integration/defs/perf/pytorch_model_config.py index 2f08739bb2fc..7a68b96e1151 100644 --- a/tests/integration/defs/perf/pytorch_model_config.py +++ b/tests/integration/defs/perf/pytorch_model_config.py @@ -17,8 +17,6 @@ Model pytorch yaml config for trtllm-bench perf tests """ -from tensorrt_llm.llmapi import KvCacheConfig - def recursive_update(d, u): for k, v in u.items(): @@ -202,9 +200,10 @@ def get_model_yaml_config(model_label: str, lora_config['lora_config']['max_lora_rank'] = 320 base_config.update(lora_config) - kv_cache_config = base_config.get('kv_cache_config', KvCacheConfig()) + kv_cache_config = base_config.get('kv_cache_config', {}) if 'kv_cache_dtype' in base_config: - kv_cache_config.dtype = base_config.pop('kv_cache_dtype', 'auto') + kv_cache_dtype = base_config.pop('kv_cache_dtype', 'auto') + kv_cache_config['dtype'] = kv_cache_dtype base_config.update({'kv_cache_config': kv_cache_config}) return base_config