"""Local-HF chat adapter matching the original host defaults.""" from contextlib import redirect_stdout import json,sys,time def chat(data): import torch from transformers import AutoModelForCausalLM,AutoTokenizer,StoppingCriteria,StoppingCriteriaList tokenizer=AutoTokenizer.from_pretrained(data['model_path'],local_files_only=True,trust_remote_code=False) model=AutoModelForCausalLM.from_pretrained(data['model_path'],torch_dtype='auto',device_map={'':'cuda:0'},low_cpu_mem_usage=True,local_files_only=True,trust_remote_code=False) model.eval() if data.get('seed') is not None:torch.manual_seed(data['seed']) prompt=tokenizer.apply_chat_template(data['messages'],tokenize=False,add_generation_prompt=True) encoded=tokenizer(prompt,return_tensors='pt',truncation=True,max_length=4096).to(model.device) count=encoded['input_ids'].shape[-1] deadline=time.monotonic()+max(1,data.get('generation_seconds',250)) class StopAtDeadline(StoppingCriteria): def __call__(self,input_ids,scores,**kwargs):return time.monotonic()>=deadline kwargs={'max_new_tokens':data['max_tokens'],'do_sample':data['temperature']>0,'repetition_penalty':data['repetition_penalty'], 'stopping_criteria':StoppingCriteriaList([StopAtDeadline()]),'pad_token_id':tokenizer.pad_token_id or tokenizer.eos_token_id} if kwargs['do_sample']:kwargs.update(temperature=data['temperature'],top_p=data['top_p']) with torch.inference_mode():tokens=model.generate(**encoded,**kwargs) generated=tokens[0,count:] return {'text':tokenizer.decode(generated,skip_special_tokens=True),'usage':{'prompt_tokens':count,'completion_tokens':len(generated)}, 'summary':{'backend':'transformers','finish_reason':'time_limit' if time.monotonic()>=deadline else 'completed'},'files':[]} if __name__=='__main__': try: with redirect_stdout(sys.stderr):result=chat(json.load(sys.stdin)) json.dump(result,sys.stdout,ensure_ascii=False) except Exception as error: json.dump({'error':str(error)[:1600]},sys.stdout,ensure_ascii=False) raise SystemExit(1)