Gemma Support (#256)

This commit is contained in:
Liangsheng Yin
2024-03-11 12:14:27 +08:00
committed by GitHub
parent 64fe311593
commit 89885b31ef
10 changed files with 428 additions and 55 deletions
+7 -5
View File
@@ -1,7 +1,5 @@
import os
from typing import Optional, Union
from typing import Optional
import torch
from sglang.srt.hf_transformers_utils import get_config, get_context_length
@@ -17,14 +15,18 @@ class ModelConfig:
self.trust_remote_code = trust_remote_code
self.revision = revision
self.hf_config = get_config(self.path, trust_remote_code, revision)
if context_length is not None:
self.context_len = context_length
else:
self.context_len = get_context_length(self.hf_config)
# Unify the config keys for hf_config
self.head_dim = self.hf_config.hidden_size // self.hf_config.num_attention_heads
self.head_dim = getattr(
self.hf_config,
"head_dim",
self.hf_config.hidden_size // self.hf_config.num_attention_heads,
)
self.num_attention_heads = self.hf_config.num_attention_heads
self.num_key_value_heads = getattr(self.hf_config, "num_key_value_heads", None)
if self.num_key_value_heads is None: