Download adapt_single_user_lambda.py from hulehule/pllm2-full-dump: direct link, hf CLI and curl.
- Browser
- Download file 11.9 kB
-
https://huggingface.co/hulehule/pllm2-full-dump/resolve/main/adapt_single_user_lambda.py
- Command line
-
hf download hf://hulehule/pllm2-full-dump/adapt_single_user_lambda.py
-
curl -L -o adapt_single_user_lambda.py https://huggingface.co/hulehule/pllm2-full-dump/resolve/main/adapt_single_user_lambda.py
11.9 kB
| #!/usr/bin/env python3 | |
| # -*- coding: utf-8 -*- | |
| """ | |
| inference_lamp4_cluster_lora.py | |
| Run inference for LaMP-4 headline generation using | |
| cluster-specific LoRA adapters trained with your LoRATrainer script. | |
| - Input: | |
| - LaMP-4 questions JSON (train/dev/test, user-based) | |
| - LaMP-4 outputs JSON (gold headlines) | |
| - cluster_info JSON (e.g., cluster_info_k60_kmeans.json) | |
| - one cluster_id and its LoRA dir (e.g., cluster_loras/cluster_12) | |
| - Output: | |
| - A JSON file with predictions for all samples belonging | |
| to users in the given cluster. | |
| """ | |
| import json | |
| import os | |
| import argparse | |
| from typing import List, Dict, Any, Tuple | |
| import torch | |
| from tqdm import tqdm | |
| from transformers import AutoTokenizer, AutoModelForCausalLM | |
| # ---------- Prompt generator (same logic as training) ---------- | |
| class LongLaMPPromptGenerator: | |
| def __init__(self, task_type: str = "lamp4_headline", max_length: int = 512, tokenizer=None): | |
| self.task_type = task_type | |
| self.max_length = max_length | |
| self.tokenizer = tokenizer | |
| def create_generation_news_prompt( | |
| self, | |
| inp: str, | |
| profile_data: List[Dict[str, Any]], | |
| max_length: int = None, | |
| tokenizer=None, | |
| ) -> str: | |
| """ | |
| LaMP-4 Headline Generation prompt. | |
| Mirrors the training-time prompt in your LoRATrainer. | |
| """ | |
| if not profile_data: | |
| return f'Generate a headline for the following article "{inp}". OUTPUT:' | |
| prompts = [] | |
| for p in profile_data: | |
| if "title" in p and "text" in p: | |
| text = p["text"] | |
| prompt = f'"{p["title"]}" is the title for "{text}"' | |
| prompts.append(prompt) | |
| if prompts: | |
| history_text = ", and ".join(prompts) | |
| return ( | |
| f'{history_text}. ' | |
| f'Following these examples, generate a headline for the following article: "{inp}". OUTPUT:' | |
| ) | |
| return f'Generate a headline for the following article "{inp}". OUTPUT:' | |
| def generate_prompt( | |
| self, | |
| input_text: str, | |
| profile_data: List[Dict[str, Any]], | |
| task_type: str = None, | |
| max_length: int = None, | |
| ) -> str: | |
| task = task_type or self.task_type | |
| if task != "lamp4_headline": | |
| raise ValueError(f"This script only supports lamp4_headline, got {task}") | |
| return self.create_generation_news_prompt(input_text, profile_data, max_length or self.max_length, self.tokenizer) | |
| # ---------- Data loading: mirror LaMP-4 conversion ---------- | |
| def detect_lamp4_outputs(raw_outputs: Any) -> List[Dict[str, Any]]: | |
| if isinstance(raw_outputs, dict): | |
| if "golds" in raw_outputs: | |
| outputs_list = raw_outputs["golds"] | |
| else: | |
| raise ValueError(f"Outputs dict has no 'golds' field. Keys: {list(raw_outputs.keys())}") | |
| elif isinstance(raw_outputs, list): | |
| outputs_list = raw_outputs | |
| else: | |
| raise ValueError(f"Unknown outputs format: {type(raw_outputs)}") | |
| if not outputs_list: | |
| raise ValueError("outputs_list is empty") | |
| first = outputs_list[0] | |
| if "id" not in first or "output" not in first: | |
| raise ValueError(f"Outputs element missing 'id' or 'output': keys={list(first.keys())}") | |
| return outputs_list | |
| def convert_lamp4_data_for_eval( | |
| questions_data: List[Dict[str, Any]], | |
| raw_outputs_data: Any, | |
| ) -> List[Dict[str, Any]]: | |
| """ | |
| Convert LaMP-4 user-based questions + outputs into a list | |
| of per-sample dicts with: | |
| { | |
| "user_id": str, | |
| "article": str, | |
| "gold": str, | |
| "profile_data": List[{title, text}], | |
| "sample_id": original sample id (optional) | |
| } | |
| """ | |
| outputs_data = detect_lamp4_outputs(raw_outputs_data) | |
| id_to_output: Dict[str, str] = {} | |
| for item in outputs_data: | |
| ex_id = str(item["id"]) | |
| id_to_output[ex_id] = item["output"] | |
| samples: List[Dict[str, Any]] = [] | |
| skipped_no_output = 0 | |
| for item in questions_data: | |
| user_id = str(item.get("id", "unknown")) | |
| if user_id not in id_to_output: | |
| skipped_no_output += 1 | |
| continue | |
| article = item.get("input", "") | |
| gold = id_to_output[user_id] | |
| if not article or not gold: | |
| continue | |
| profile_data: List[Dict[str, Any]] = [] | |
| profile = item.get("profile", []) | |
| for article_obj in profile[:10]: | |
| title = article_obj.get("title", "") | |
| text = article_obj.get("text", "") | |
| if title and text: | |
| if len(text) > 1000: | |
| text = text[:1000] | |
| profile_data.append({"title": title, "text": text}) | |
| samples.append( | |
| { | |
| "user_id": user_id, | |
| "article": article, | |
| "gold": gold, | |
| "profile_data": profile_data, | |
| "sample_id": user_id, # reuse user id as sample id in this setup | |
| } | |
| ) | |
| print(f"✅ Converted LaMP-4 for eval: {len(samples)} samples") | |
| if skipped_no_output > 0: | |
| print(f"⚠️ Skipped {skipped_no_output} questions with no matching output") | |
| return samples | |
| def load_lamp4_eval_data(questions_path: str, outputs_path: str) -> List[Dict[str, Any]]: | |
| print(f"Loading LaMP-4 questions from: {questions_path}") | |
| with open(questions_path, "r", encoding="utf-8") as f: | |
| raw_q = json.load(f) | |
| if isinstance(raw_q, dict) and "questions" in raw_q: | |
| questions = raw_q["questions"] | |
| elif isinstance(raw_q, list): | |
| questions = raw_q | |
| else: | |
| raise ValueError(f"Unknown questions format: {type(raw_q)}") | |
| print(f"Loading LaMP-4 outputs from: {outputs_path}") | |
| with open(outputs_path, "r", encoding="utf-8") as f: | |
| raw_out = json.load(f) | |
| samples = convert_lamp4_data_for_eval(questions, raw_out) | |
| return samples | |
| # ---------- Cluster info ---------- | |
| def load_cluster_info(cluster_info_path: str) -> Dict[str, Any]: | |
| print(f"Loading cluster info from: {cluster_info_path}") | |
| with open(cluster_info_path, "r", encoding="utf-8") as f: | |
| info = json.load(f) | |
| if "clusters" not in info: | |
| raise ValueError("cluster_info missing 'clusters' field") | |
| clusters = info["clusters"] # {cluster_id: [user_ids]} | |
| # normalize keys to str | |
| clusters = {str(k): v for k, v in clusters.items()} | |
| info["clusters"] = clusters | |
| print(f"✅ Loaded cluster_info: {info.get('n_clusters', len(clusters))} clusters") | |
| return info | |
| # ---------- Inference ---------- | |
| def generate_headline( | |
| model, | |
| tokenizer, | |
| prompt: str, | |
| max_new_tokens: int = 32, | |
| temperature: float = 0.0, | |
| top_p: float = 1.0, | |
| ) -> str: | |
| inputs = tokenizer(prompt, return_tensors="pt").to(model.device) | |
| with torch.no_grad(): | |
| output_ids = model.generate( | |
| **inputs, | |
| max_new_tokens=max_new_tokens, | |
| do_sample=(temperature > 0), | |
| temperature=temperature if temperature > 0 else 1.0, | |
| top_p=top_p, | |
| eos_token_id=tokenizer.eos_token_id, | |
| pad_token_id=tokenizer.pad_token_id, | |
| ) | |
| # take only the generated part after the prompt | |
| gen_ids = output_ids[0][inputs["input_ids"].shape[1]:] | |
| text = tokenizer.decode(gen_ids, skip_special_tokens=True) | |
| return text.strip() | |
| def run_inference_for_cluster( | |
| model_name: str, | |
| lora_dir: str, | |
| questions_path: str, | |
| outputs_path: str, | |
| cluster_info_path: str, | |
| cluster_id: str, | |
| output_file: str, | |
| max_new_tokens: int = 32, | |
| ): | |
| """ | |
| Run inference for one cluster: | |
| - Load base+LoRA from lora_dir | |
| - Filter samples whose user_id is in cluster_users | |
| - Generate headlines and save to JSON | |
| """ | |
| # 1) Load eval data | |
| all_samples = load_lamp4_eval_data(questions_path, outputs_path) | |
| # 2) Load cluster info | |
| cluster_info = load_cluster_info(cluster_info_path) | |
| clusters = cluster_info["clusters"] | |
| if cluster_id not in clusters: | |
| raise ValueError(f"Cluster id {cluster_id} not found in cluster_info") | |
| cluster_users = set(clusters[cluster_id]) | |
| print(f"Cluster {cluster_id}: {len(cluster_users)} users") | |
| # 3) Filter samples for this cluster | |
| cluster_samples = [s for s in all_samples if s["user_id"] in cluster_users] | |
| print(f"Cluster {cluster_id}: {len(cluster_samples)} samples to evaluate") | |
| if not cluster_samples: | |
| print("No samples for this cluster, abort.") | |
| return | |
| # 4) Load model + tokenizer from LoRA dir | |
| print(f"Loading LoRA model from: {lora_dir}") | |
| model = AutoModelForCausalLM.from_pretrained( | |
| lora_dir, | |
| torch_dtype=torch.bfloat16, | |
| device_map="auto", | |
| trust_remote_code=True, | |
| ) | |
| tokenizer = AutoTokenizer.from_pretrained(model_name, trust_remote_code=True) | |
| if tokenizer.pad_token is None: | |
| tokenizer.pad_token = tokenizer.eos_token | |
| tokenizer.padding_side = "right" | |
| tokenizer.truncation_side = "right" | |
| prompt_generator = LongLaMPPromptGenerator(task_type="lamp4_headline", tokenizer=tokenizer) | |
| # 5) Run generation | |
| results = [] | |
| for sample in tqdm(cluster_samples, desc=f"Cluster {cluster_id} inference"): | |
| prompt = prompt_generator.generate_prompt( | |
| input_text=sample["article"], | |
| profile_data=sample["profile_data"], | |
| task_type="lamp4_headline", | |
| ) | |
| pred = generate_headline( | |
| model=model, | |
| tokenizer=tokenizer, | |
| prompt=prompt, | |
| max_new_tokens=max_new_tokens, | |
| temperature=0.0, | |
| top_p=1.0, | |
| ) | |
| results.append( | |
| { | |
| "user_id": sample["user_id"], | |
| "sample_id": sample["sample_id"], | |
| "article": sample["article"], | |
| "gold": sample["gold"], | |
| "prompt": prompt, | |
| "prediction": pred, | |
| "cluster_id": cluster_id, | |
| } | |
| ) | |
| # 6) Save results | |
| os.makedirs(os.path.dirname(output_file), exist_ok=True) | |
| with open(output_file, "w", encoding="utf-8") as f: | |
| json.dump(results, f, indent=2, ensure_ascii=False) | |
| print(f"✅ Saved {len(results)} predictions to {output_file}") | |
| def main(): | |
| parser = argparse.ArgumentParser(description="Inference for LaMP-4 cluster-specific LoRA.") | |
| parser.add_argument("--model_name", type=str, default="Qwen/Qwen2.5-7B-Instruct") | |
| parser.add_argument("--lora_dir", type=str, required=True, | |
| help="Directory of the trained LoRA for this cluster (e.g. cluster_loras/cluster_12)") | |
| parser.add_argument("--questions_json", type=str, required=True, | |
| help="LaMP-4 questions JSON (user-based)") | |
| parser.add_argument("--outputs_json", type=str, required=True, | |
| help="LaMP-4 outputs JSON (golds)") | |
| parser.add_argument("--cluster_info_path", type=str, required=True, | |
| help="cluster_info JSON (e.g., cluster_info_k60_kmeans.json)") | |
| parser.add_argument("--cluster_id", type=str, required=True, | |
| help="Cluster ID to evaluate, e.g. '12'") | |
| parser.add_argument("--output_file", type=str, required=True, | |
| help="Path to save predictions JSON") | |
| parser.add_argument("--max_new_tokens", type=int, default=32) | |
| args = parser.parse_args() | |
| run_inference_for_cluster( | |
| model_name=args.model_name, | |
| lora_dir=args.lora_dir, | |
| questions_path=args.questions_json, | |
| outputs_path=args.outputs_json, | |
| cluster_info_path=args.cluster_info_path, | |
| cluster_id=args.cluster_id, | |
| output_file=args.output_file, | |
| max_new_tokens=args.max_new_tokens, | |
| ) | |
| if __name__ == "__main__": | |
| main() | |