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()