lspcloud/prolific-preferences-personalized
0
1"""Tinker inference client. Supports both base models and fine-tuned checkpoints."""2import re3 4import streamlit as st5 6 7@st.cache_resource8def _get_tinker_clients(model_name: str, sampler_path: str = ""):9 """10 Initialise and cache the Tinker sampling client, renderer, and tokenizer.11 If sampler_path is provided, loads from that checkpoint (fine-tuned model).12 Otherwise, loads the base model_name.13 Cache key includes both so different variants get different clients.14 """15 import tinker16 from tinker import types as tinker_types17 from tinker_cookbook import renderers18 from tinker_cookbook.model_info import get_recommended_renderer_name19 from tinker_cookbook.tokenizer_utils import get_tokenizer20 21 service_client = tinker.ServiceClient()22 if sampler_path:23 print(f"[MODEL] Loading fine-tuned checkpoint: {sampler_path}")24 sampling_client = service_client.create_sampling_client(model_path=sampler_path)25 else:26 print(f"[MODEL] Loading base model: {model_name}")27 sampling_client = service_client.create_sampling_client(base_model=model_name)28 29 tokenizer = get_tokenizer(model_name)30 renderer_name = get_recommended_renderer_name(model_name)31 renderer = renderers.get_renderer(renderer_name, tokenizer)32 return sampling_client, renderer, tinker_types33 34 35def call_model(messages: list, cfg: dict) -> str:36 """Send a message list to Tinker and return cleaned response text."""37 model_name = cfg["model_name"]38 sampler_path = cfg.get("sampler_path", "")39 print(f"[MODEL] model_name={model_name} sampler_path={sampler_path or '(base)'}")40 print(f"[MODEL] num_messages={len(messages)}")41 print(f"[MODEL] roles={[m['role'] for m in messages]}")42 if messages:43 print(f"[MODEL] system_prompt[:150]={messages[0]['content'][:150]}")44 45 try:46 from tinker_cookbook import renderers as tinker_renderers47 48 sampling_client, renderer, tinker_types = _get_tinker_clients(model_name, sampler_path)49 50 prompt = renderer.build_generation_prompt(messages)51 params = tinker_types.SamplingParams(52 max_tokens=1000,53 temperature=0.7,54 stop=renderer.get_stop_sequences(),55 )56 result = sampling_client.sample(57 prompt=prompt,58 sampling_params=params,59 num_samples=1,60 ).result()61 62 parsed_message, _ = renderer.parse_response(result.sequences[0].tokens)63 content = tinker_renderers.format_content_as_string(parsed_message["content"])64 65 content = re.sub(r"<think>.*?</think>", "", content, flags=re.DOTALL).strip()66 content = re.sub(r"<\|[^|]*\|>", "", content).strip()67 match = re.search(r"(.{40,}?)\1{4,}", content, flags=re.DOTALL)68 if match:69 content = content[: match.start() + len(match.group(1))].strip()70 if not content or len(content.split()) < 3:71 raise ValueError("Model output cleanup yielded no usable content.")72 73 return content74 75 except Exception as e:76 print(f"[MODEL] Tinker error: {e}")77 return f"[Model error: {e}]"