Gemma Support (#256)
This commit is contained in:
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user