from __future__ import annotations import argparse from pathlib import Path import torch from transformers import AutoTokenizer from dot_rd.config import ModelConfig from dot_rd.export import load_exported_core, load_inference_checkpoint from dot_rd.model import DotRecurrentDepthModel def _token_ids(value: object) -> list[int]: if hasattr(value, "input_ids"): value = value.input_ids if isinstance(value, torch.Tensor): value = value.tolist() if isinstance(value, list) and value and isinstance(value[0], list): value = value[0] if not isinstance(value, list) or not all(isinstance(token, int) for token in value): raise TypeError("chat template did not return a token id list") return value @torch.inference_mode() def generate( model: DotRecurrentDepthModel, tokenizer: object, prompt: str, *, max_new_tokens: int, max_context_tokens: int, ) -> str: messages = [ {"role": "system", "content": "You are Dot, a local reasoning model."}, {"role": "user", "content": prompt}, ] prompt_ids = _token_ids( tokenizer.apply_chat_template( messages, tokenize=True, add_generation_prompt=True, enable_thinking=True, ) )[-max_context_tokens:] input_ids = torch.tensor([prompt_ids], dtype=torch.long, device="cuda") attention_mask = torch.ones_like(input_ids) generated: list[int] = [] eos_token_id = tokenizer.eos_token_id if eos_token_id is None: raise ValueError("Dot tokenizer has no EOS token") for _ in range(max_new_tokens): output = model( input_ids=input_ids, attention_mask=attention_mask, use_cache=False, logits_to_keep=1, ) next_token = int(output.logits[:, -1].argmax(dim=-1).item()) if next_token == eos_token_id: break generated.append(next_token) input_ids = torch.cat( (input_ids, torch.tensor([[next_token]], device=input_ids.device)), dim=1 ) attention_mask = torch.cat( (attention_mask, torch.ones((1, 1), dtype=torch.long, device=input_ids.device)), dim=1, ) suffix = tokenizer.decode(generated, skip_special_tokens=True).strip() return suffix if suffix.startswith("") else f"\n{suffix}" def main() -> None: parser = argparse.ArgumentParser(description="Run Dot v0.4 with greedy decoding") parser.add_argument("--model", default=".", help="local Dot repository path") parser.add_argument( "--checkpoint", help="optional Dot inference-checkpoint directory containing manifest.json", ) parser.add_argument("--prompt", required=True) parser.add_argument("--max-new-tokens", type=int, default=256) parser.add_argument("--max-context-tokens", type=int, default=4096) args = parser.parse_args() if not torch.cuda.is_available(): raise RuntimeError("the verified Dot v0.4 runtime requires CUDA") model_path = str(Path(args.model).resolve()) tokenizer = AutoTokenizer.from_pretrained(model_path) config = ModelConfig( base_model=model_path, insertion_after=15, source_layers=(12, 13, 14, 15), max_loops=8, active_loops=4, initial_loop_scale=0.01, attention_implementation="sdpa", ) model = DotRecurrentDepthModel.from_pretrained( config, dtype=torch.bfloat16, device_map=None, ).to("cuda").eval() manifest = ( load_inference_checkpoint(args.checkpoint, model) if args.checkpoint else load_exported_core(model_path, model) ) print( generate( model, tokenizer, args.prompt, max_new_tokens=args.max_new_tokens, max_context_tokens=args.max_context_tokens, ) ) print(f"\n[Dot release step {manifest['source_step']}]") if __name__ == "__main__": main()