aabbdev's picture
Publish RWKV7-1.5B-SMI-20260822
6d93aeb verified
Raw
History Blame Contribute Delete
6.45 kB
from __future__ import annotations
import json
from pathlib import Path
import torch
from huggingface_hub import snapshot_download
from safetensors.torch import load_file
from transformers import AutoTokenizer, PreTrainedConfig
from .runtime import RWKV7Config, RWKV7ForCausalLM
SQUEEZE_MARKERS = (
".x_",
".k_",
"att.r",
"att.w",
"att.v0",
"att.v1",
"att.v2",
"att.a",
"att.g",
)
def _resolve_model(model: str) -> tuple[Path, bool]:
local = Path(model).expanduser()
if local.is_dir():
return local.resolve(), False
return Path(snapshot_download(model)), True
def _optimized_config(native: dict, backend: str, state_dtype: str) -> RWKV7Config:
return RWKV7Config(
vocab_size=native["vocab_size"],
hidden_size=native["hidden_size"],
num_hidden_layers=native["num_hidden_layers"],
head_size=native["head_dim"],
intermediate_size=native["intermediate_size"],
decay_lora_rank=native["decay_low_rank_dim"],
a_lora_rank=native["a_low_rank_dim"],
gate_lora_rank=native["gate_low_rank_dim"],
value_lora_rank=native["v_low_rank_dim"],
layer_norm_epsilon=native.get("norm_eps", 1e-5),
use_cache=native.get("use_cache", True),
kernel_backend=backend,
recurrent_state_dtype=state_dtype,
tie_word_embeddings=native.get("tie_word_embeddings", False),
bos_token_id=native.get("bos_token_id"),
eos_token_id=native.get("eos_token_id", 0),
pad_token_id=native.get("pad_token_id", 0),
)
def _native_key_to_optimized(key: str) -> str:
return key.removeprefix("rwkv7.")
def _native_tensor_to_optimized(key: str, tensor: torch.Tensor) -> torch.Tensor:
return tensor.squeeze() if any(marker in key for marker in SQUEEZE_MARKERS) else tensor
def _checked_weight(path: Path, *, hub_snapshot: bool) -> Path:
if path.is_symlink():
if not hub_snapshot:
raise RuntimeError(f"local safetensor must not be a symlink: {path.name}")
path = path.resolve(strict=True)
if not path.is_file():
raise RuntimeError(f"safetensor must be a regular file: {path.name}")
return path
def _weight_plan(
model_dir: Path, *, hub_snapshot: bool
) -> list[tuple[Path, set[str] | None]]:
present = sorted(model_dir.glob("model*.safetensors"))
if not present:
raise FileNotFoundError(f"no safetensors found in {model_dir}")
checked = {
path.name: _checked_weight(path, hub_snapshot=hub_snapshot) for path in present
}
index_path = model_dir / "model.safetensors.index.json"
if not index_path.is_file():
if [path.name for path in present] != ["model.safetensors"]:
raise RuntimeError("multiple safetensors require model.safetensors.index.json")
return [(checked["model.safetensors"], None)]
index = json.loads(index_path.read_text(encoding="utf-8"))
weight_map = index.get("weight_map")
if not isinstance(weight_map, dict) or not weight_map:
raise RuntimeError("model.safetensors.index.json has no weight_map")
keys_by_file: dict[str, set[str]] = {}
for key, filename in weight_map.items():
if not isinstance(key, str) or not isinstance(filename, str):
raise RuntimeError("invalid weight_map entry")
keys_by_file.setdefault(filename, set()).add(key)
if set(keys_by_file) != {path.name for path in present}:
raise RuntimeError("indexed and present safetensors files differ")
return [(checked[filename], keys_by_file[filename]) for filename in sorted(keys_by_file)]
def _load_native_model(
model_dir: Path,
native_config: dict,
backend: str,
state_dtype: str,
*,
hub_snapshot: bool,
) -> RWKV7ForCausalLM:
config = _optimized_config(native_config, backend, state_dtype)
with torch.device("meta"):
model = RWKV7ForCausalLM(config)
expected = set(model.state_dict())
seen: set[str] = set()
for weight_file, indexed_keys in _weight_plan(
model_dir, hub_snapshot=hub_snapshot
):
native_shard = load_file(weight_file, device="cpu")
if indexed_keys is not None and set(native_shard) != indexed_keys:
raise RuntimeError(f"tensor keys in {weight_file.name} do not match the index")
shard = {}
for native_key, tensor in native_shard.items():
optimized_key = _native_key_to_optimized(native_key)
if optimized_key in seen or optimized_key in shard:
raise RuntimeError(f"duplicate optimized tensor key: {optimized_key}")
shard[optimized_key] = _native_tensor_to_optimized(optimized_key, tensor)
unexpected = sorted(set(shard) - expected)
if unexpected:
raise RuntimeError(f"native checkpoint has unexpected keys: {unexpected}")
model.load_state_dict(shard, strict=False, assign=True)
seen.update(shard)
missing = sorted(expected - seen)
if missing:
raise RuntimeError(f"native checkpoint is missing optimized runtime keys: {missing}")
return model
def _load_tokenizer(model_dir: Path):
return AutoTokenizer.from_pretrained(
model_dir,
config=PreTrainedConfig(),
local_files_only=True,
)
def load_model_and_tokenizer(
model: str,
*,
device: str,
dtype: torch.dtype | None,
backend: str,
state_dtype: str,
):
model_dir, hub_snapshot = _resolve_model(model)
native_config = json.loads((model_dir / "config.json").read_text(encoding="utf-8"))
architectures = set(native_config.get("architectures", []))
if architectures != {"Rwkv7ForCausalLM"}:
raise ValueError(f"unsupported RWKV-7 architecture: {sorted(architectures)}")
loaded = _load_native_model(
model_dir,
native_config,
backend,
state_dtype,
hub_snapshot=hub_snapshot,
)
if dtype is None:
dtype_name = str(native_config.get("dtype", "bfloat16")).removeprefix("torch.")
try:
dtype = {
"bfloat16": torch.bfloat16,
"float16": torch.float16,
"float32": torch.float32,
}[dtype_name]
except KeyError as error:
raise ValueError(f"unsupported model dtype: {dtype_name}") from error
loaded = loaded.to(device=device, dtype=dtype).eval()
return loaded, _load_tokenizer(model_dir)