import torch from peft import PeftModel from transformers import AutoTokenizer, AutoModelForCausalLM ADAPTER_ID = 'mario-rc/emotional-rlaif-dpo-gemma-2-9b-it' BASE_ID = 'google/gemma-2-9b-it' # Pin the base revision checked when this release was prepared. BASE_REVISION = '11c9b309abf73637e4b6f9a3fa1e92e615547819' # Set this to a commit hash from this adapter's Files and versions tab for a pinned run. ADAPTER_REVISION = "main" SYSTEM_PROMPT = """You are an expert at creating dialogues. Dialogue and emotional structure: Human: (SADNESS) PROMPT. Chatbot: (SADNESS) RESPONSE_1. (HAPPINESS) RESPONSE_2. (NEUTRAL) RESPONSE_3. Dialogue rules: The response must be open-domain curated. The response should be coherent, empathetic, engaging and proactive. The chatbot RESPONSE is composed of 3 different sentences (RESPONSE_1, RESPONSE_2 and RESPONSE_3), separated by a period. Between RESPONSE_1, RESPONSE_2 and RESPONSE_3 should be a max length of 20-25 words. RESPONSE_3 must be open-ended to follow-up the conversation, so the Human is encouraged to answer with a full long sentence. Avoid yes/no questions. Emotional response rules: RESPONSE_1 must contain a SADNESS tone. RESPONSE_2 must contain a HAPPINESS tone. RESPONSE_3 must contain a NEUTRAL tone. Answer in a single turn to Human. Follow exactly the emotional structure and the emotional and dialogue rules.""" def main(): tokenizer = AutoTokenizer.from_pretrained( ADAPTER_ID, revision=ADAPTER_REVISION, trust_remote_code=False, ) base = AutoModelForCausalLM.from_pretrained( BASE_ID, revision=BASE_REVISION, trust_remote_code=False, torch_dtype=torch.bfloat16, device_map="auto", ) model = PeftModel.from_pretrained(base, ADAPTER_ID, revision=ADAPTER_REVISION) model.eval() messages = [ {"role": "system", "content": SYSTEM_PROMPT}, {"role": "user", "content": "(SADNESS) I feel overwhelmed by my exams."}, ] prompt = tokenizer.apply_chat_template( messages, tokenize=False, add_generation_prompt=True, ) inputs = tokenizer(prompt, add_special_tokens=False, return_tensors="pt") inputs = {k: v.to(model.device) for k, v in inputs.items()} stop_ids = model.generation_config.eos_token_id stop_ids = list(stop_ids) if isinstance(stop_ids, (list, tuple)) else [stop_ids] stop_ids = sorted({i for i in stop_ids + [tokenizer.eos_token_id] if i is not None}) # Include turn-ending tokens, which are not EOS in every base tokenizer. for token in ['']: token_id = tokenizer.convert_tokens_to_ids(token) if token_id is not None and token_id != tokenizer.unk_token_id and token_id not in stop_ids: stop_ids.append(token_id) pad_id = tokenizer.pad_token_id if tokenizer.pad_token_id is not None else stop_ids[0] with torch.inference_mode(): output = model.generate( **inputs, max_new_tokens=128, do_sample=False, eos_token_id=stop_ids, pad_token_id=pad_id, ) answer = tokenizer.decode(output[0, inputs["input_ids"].shape[-1]:], skip_special_tokens=True) print(answer) if __name__ == "__main__": main()