Fix added tokens config with sensible filter (#17905)
This commit is contained in:
@@ -128,8 +128,8 @@ class LoRAAdapter(nn.Module):
|
|||||||
# added/extra token emb
|
# added/extra token emb
|
||||||
self.added_tokens_embeddings[name] = loaded_weight.cpu()
|
self.added_tokens_embeddings[name] = loaded_weight.cpu()
|
||||||
assert loaded_weight.shape[0] == self.config.lora_added_tokens_size, (
|
assert loaded_weight.shape[0] == self.config.lora_added_tokens_size, (
|
||||||
f"LoRA adapter {self.uid} has extra_vocab_size {self.config.extra_vocab_size} specified in the config, "
|
f"LoRA adapter {self.uid} has lora_added_tokens_size {self.config.lora_added_tokens_size} specified in the config, "
|
||||||
f"but the loaded weight has {loaded_weight.shape[0]} extra vocab size"
|
f"but the loaded weight '{name}' has shape {loaded_weight.shape[0]} in first dimension"
|
||||||
)
|
)
|
||||||
|
|
||||||
def _normalize_weights(self):
|
def _normalize_weights(self):
|
||||||
|
|||||||
@@ -13,11 +13,14 @@
|
|||||||
# ==============================================================================
|
# ==============================================================================
|
||||||
|
|
||||||
import json
|
import json
|
||||||
|
import logging
|
||||||
import os
|
import os
|
||||||
from typing import Dict, Optional
|
from typing import Dict, Optional
|
||||||
|
|
||||||
from huggingface_hub import snapshot_download
|
from huggingface_hub import snapshot_download
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
class LoRAConfig:
|
class LoRAConfig:
|
||||||
def __init__(
|
def __init__(
|
||||||
@@ -25,6 +28,7 @@ class LoRAConfig:
|
|||||||
path: Optional[str] = None,
|
path: Optional[str] = None,
|
||||||
config_dict: Optional[Dict] = None,
|
config_dict: Optional[Dict] = None,
|
||||||
added_tokens_config: Optional[Dict] = None,
|
added_tokens_config: Optional[Dict] = None,
|
||||||
|
base_vocab_size: Optional[int] = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
self.path = path
|
self.path = path
|
||||||
|
|
||||||
@@ -38,17 +42,41 @@ class LoRAConfig:
|
|||||||
self.target_modules = self.hf_config["target_modules"]
|
self.target_modules = self.hf_config["target_modules"]
|
||||||
self.r = self.hf_config["r"]
|
self.r = self.hf_config["r"]
|
||||||
self.lora_alpha = self.hf_config["lora_alpha"]
|
self.lora_alpha = self.hf_config["lora_alpha"]
|
||||||
|
|
||||||
|
# Filter fake added tokens: tokens with ID < base_vocab_size are already
|
||||||
|
# part of the base vocabulary and should not be treated as added tokens.
|
||||||
|
# This commonly happens when added_tokens.json is copied from the base
|
||||||
|
# model's tokenizer.
|
||||||
|
if self.added_tokens_config and base_vocab_size is not None:
|
||||||
|
self.added_tokens_config = {
|
||||||
|
token: token_id
|
||||||
|
for token, token_id in self.added_tokens_config.items()
|
||||||
|
if token_id >= base_vocab_size
|
||||||
|
}
|
||||||
|
|
||||||
self.lora_added_tokens_size = (
|
self.lora_added_tokens_size = (
|
||||||
len(self.added_tokens_config) if self.added_tokens_config is not None else 0
|
len(self.added_tokens_config) if self.added_tokens_config is not None else 0
|
||||||
)
|
)
|
||||||
|
|
||||||
|
if self.lora_added_tokens_size > 0:
|
||||||
|
raise ValueError(
|
||||||
|
f"LoRA adapter has {self.lora_added_tokens_size} added tokens, "
|
||||||
|
f"but added tokens are not supported yet. "
|
||||||
|
f"Added tokens: {self.added_tokens_config}"
|
||||||
|
)
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def from_dict(
|
def from_dict(
|
||||||
cls,
|
cls,
|
||||||
config_dict: Dict,
|
config_dict: Dict,
|
||||||
added_tokens_config: Optional[Dict] = None,
|
added_tokens_config: Optional[Dict] = None,
|
||||||
|
base_vocab_size: Optional[int] = None,
|
||||||
) -> "LoRAConfig":
|
) -> "LoRAConfig":
|
||||||
return cls(config_dict=config_dict, added_tokens_config=added_tokens_config)
|
return cls(
|
||||||
|
config_dict=config_dict,
|
||||||
|
added_tokens_config=added_tokens_config,
|
||||||
|
base_vocab_size=base_vocab_size,
|
||||||
|
)
|
||||||
|
|
||||||
def get_lora_config(self, dummy=False):
|
def get_lora_config(self, dummy=False):
|
||||||
if dummy:
|
if dummy:
|
||||||
@@ -82,9 +110,5 @@ class LoRAConfig:
|
|||||||
with open(added_tokens_path, "r") as f:
|
with open(added_tokens_path, "r") as f:
|
||||||
return json.load(f)
|
return json.load(f)
|
||||||
except json.JSONDecodeError as e:
|
except json.JSONDecodeError as e:
|
||||||
# Log warning but don't crash if JSON is malformed
|
|
||||||
import logging
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
|
||||||
logger.warning(f"Failed to parse added_tokens.json: {e}")
|
logger.warning(f"Failed to parse added_tokens.json: {e}")
|
||||||
return None
|
return None
|
||||||
|
|||||||
@@ -136,7 +136,10 @@ class LoRAManager:
|
|||||||
|
|
||||||
try:
|
try:
|
||||||
# load configs
|
# load configs
|
||||||
new_adapter = LoRAConfig(lora_ref.lora_path)
|
new_adapter = LoRAConfig(
|
||||||
|
lora_ref.lora_path,
|
||||||
|
base_vocab_size=self.base_hf_config.vocab_size,
|
||||||
|
)
|
||||||
self.validate_new_adapter(new_adapter, lora_ref)
|
self.validate_new_adapter(new_adapter, lora_ref)
|
||||||
self.configs[lora_ref.lora_id] = new_adapter
|
self.configs[lora_ref.lora_id] = new_adapter
|
||||||
|
|
||||||
@@ -602,7 +605,11 @@ class LoRAManager:
|
|||||||
), f"LoRA adapter with ID {lora_ref.lora_id} is already loaded. This should have been verified before request is sent to the backend."
|
), f"LoRA adapter with ID {lora_ref.lora_id} is already loaded. This should have been verified before request is sent to the backend."
|
||||||
|
|
||||||
try:
|
try:
|
||||||
new_adapter = LoRAConfig.from_dict(config_dict, added_tokens_config)
|
new_adapter = LoRAConfig.from_dict(
|
||||||
|
config_dict,
|
||||||
|
added_tokens_config,
|
||||||
|
base_vocab_size=self.base_hf_config.vocab_size,
|
||||||
|
)
|
||||||
self.validate_new_adapter(new_adapter, lora_ref)
|
self.validate_new_adapter(new_adapter, lora_ref)
|
||||||
self.configs[lora_ref.lora_id] = new_adapter
|
self.configs[lora_ref.lora_id] = new_adapter
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user