shoumenchougou's picture
Upload folder using huggingface_hub
d2d414f verified
Raw History Blame
7.19 kB
from __future__ import annotations
import argparse
import importlib
import os
import re
import sys
from types import ModuleType
from pathlib import Path
from typing import Any
import torch
from transformers import StoppingCriteria, StoppingCriteriaList
if __package__ in {None, ""}:
package_name = "_rwkv7_release_inference"
package = ModuleType(package_name)
package.__package__ = package_name
package.__path__ = [str(Path(__file__).resolve().parent)]
sys.modules[package_name] = package
load_model_and_tokenizer = importlib.import_module(
f"{package_name}.model_loader"
).load_model_and_tokenizer
else:
from .model_loader import load_model_and_tokenizer
DTYPES = {
"bfloat16": torch.bfloat16,
"float16": torch.float16,
"float32": torch.float32,
}
STOP_TEXT = "\n\nUser:"
class StopOnText(StoppingCriteria):
def __init__(self, tokenizer: Any, prompt_length: int, stop_text: str) -> None:
self.tokenizer = tokenizer
self.prompt_length = prompt_length
self.stop_text = stop_text
def __call__(
self,
input_ids: torch.LongTensor,
scores: torch.FloatTensor,
**kwargs: Any,
) -> bool:
del scores, kwargs
completion = self.tokenizer.decode(
input_ids[0, self.prompt_length :],
skip_special_tokens=False,
)
return self.stop_text in completion
def _prompt_ids(tokenizer: Any, messages: list[dict[str, str]], thinking: bool):
tokens = tokenizer.apply_chat_template(
messages,
tokenize=True,
add_generation_prompt=True,
thinking=thinking,
return_tensors="pt",
)
if hasattr(tokens, "input_ids"):
tokens = tokens.input_ids
elif isinstance(tokens, dict):
tokens = tokens["input_ids"]
if tokens.ndim == 1:
tokens = tokens.unsqueeze(0)
return tokens
@torch.inference_mode()
def generate_completion(
model: Any,
tokenizer: Any,
messages: list[dict[str, str]],
*,
device: str,
max_new_tokens: int,
temperature: float,
top_p: float,
thinking: bool,
) -> str:
input_ids = _prompt_ids(tokenizer, messages, thinking).to(device)
prompt_length = input_ids.shape[1]
generation: dict[str, Any] = {
"input_ids": input_ids,
"attention_mask": torch.ones_like(input_ids),
"max_new_tokens": max_new_tokens,
"do_sample": temperature > 0,
"eos_token_id": 0,
"pad_token_id": 0,
"stopping_criteria": StoppingCriteriaList(
[StopOnText(tokenizer, prompt_length, STOP_TEXT)]
),
}
if temperature > 0:
generation["temperature"] = temperature
generation["top_p"] = top_p
output = model.generate(**generation)
completion_ids = output[0, prompt_length:]
completion = tokenizer.decode(completion_ids, skip_special_tokens=True)
if STOP_TEXT in completion:
completion = completion.split(STOP_TEXT, 1)[0]
return completion.strip()
def _interactive(
model: Any,
tokenizer: Any,
args: argparse.Namespace,
) -> None:
messages: list[dict[str, str]] = []
print("RWKV-7 Goose — /clear resets the conversation, /exit quits.")
while True:
try:
prompt = input(">>> ")
except EOFError:
break
if prompt == "/exit":
break
if prompt == "/clear":
messages.clear()
continue
prompt = prompt.strip()
if not prompt:
continue
messages.append({"role": "user", "content": prompt})
completion = generate_completion(
model,
tokenizer,
messages,
device=args.device,
max_new_tokens=args.max_new_tokens,
temperature=args.temperature,
top_p=args.top_p,
thinking=args.thinking,
)
print(completion)
messages.append({"role": "assistant", "content": completion})
def _file_prompts(
model: Any,
tokenizer: Any,
args: argparse.Namespace,
) -> None:
text = Path(args.input_file).read_text(encoding="utf-8")
prompts = [prompt.strip() for prompt in re.split(r"\n\s*\n", text) if prompt.strip()]
if not prompts:
raise ValueError("input file contains no prompts")
for prompt in prompts:
completion = generate_completion(
model,
tokenizer,
[{"role": "user", "content": prompt}],
device=args.device,
max_new_tokens=args.max_new_tokens,
temperature=args.temperature,
top_p=args.top_p,
thinking=args.thinking,
)
print(f"Prompt: {prompt}")
print(f"Completion: {completion}")
print()
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description="Generate with RWKV-7 Goose")
parser.add_argument("--model", required=True, help="Hub repo ID or local model directory")
mode = parser.add_mutually_exclusive_group(required=True)
mode.add_argument("--interactive", action="store_true")
mode.add_argument("--input-file")
parser.add_argument("--device", default="cuda")
parser.add_argument("--dtype", choices=("auto", *DTYPES), default="auto")
parser.add_argument("--state-dtype", choices=DTYPES, default="float32")
parser.add_argument("--backend", choices=("auto", "torch", "tilelang"), default="auto")
parser.add_argument("--max-new-tokens", type=int, default=300)
parser.add_argument("--temperature", type=float, default=1.0)
parser.add_argument("--top-p", type=float, default=0.5)
parser.add_argument("--seed", type=int, default=33377335)
parser.add_argument("--thinking", action="store_true")
return parser.parse_args()
def main() -> None:
args = parse_args()
if int(os.getenv("WORLD_SIZE", "1")) != 1:
raise RuntimeError("the bundled runtime supports one process and one GPU")
if int(os.getenv("RANK", "0")) != 0 or int(os.getenv("LOCAL_RANK", "0")) != 0:
raise RuntimeError("RANK and LOCAL_RANK must be zero")
if args.max_new_tokens <= 0:
raise ValueError("max-new-tokens must be positive")
if args.temperature < 0:
raise ValueError("temperature must be non-negative")
if not 0 < args.top_p <= 1:
raise ValueError("top-p must be in (0, 1]")
torch.manual_seed(args.seed)
model, tokenizer = load_model_and_tokenizer(
args.model,
device=args.device,
dtype=None if args.dtype == "auto" else DTYPES[args.dtype],
backend=args.backend,
state_dtype=args.state_dtype,
)
model.set_kernel_backend(args.backend)
if args.backend == "tilelang":
model.prepare_inference_weights()
if args.interactive:
_interactive(model, tokenizer, args)
else:
_file_prompts(model, tokenizer, args)
if __name__ == "__main__":
main()