Spaces:
Sleeping
Sleeping
| """ | |
| Engine wrapper for AutoGEO — exposes clean APIs for the Streamlit UI. | |
| """ | |
| import sys | |
| import os | |
| import time | |
| from typing import Optional, List, Dict | |
| # Ensure the parent is on sys.path so autogeo imports work | |
| _parent = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) | |
| if _parent not in sys.path: | |
| sys.path.insert(0, _parent) | |
| from autogeo.rewriters.core import rewrite_document | |
| from autogeo.evaluation.generative_engine import ( | |
| generate_answer_gemini, | |
| generate_answer_gpt, | |
| generate_answer_claude, | |
| generate_answer_deepseek, | |
| generate_answer_doubao, | |
| ) | |
| from autogeo.evaluation.metrics.geo_score import ( | |
| extract_citations_new, | |
| impression_pos_count_simple, | |
| impression_word_count_simple, | |
| impression_wordpos_count_simple, | |
| ) | |
| from autogeo.utils.deepseek import call_deepseek | |
| from autogeo.utils.doubao import call_doubao | |
| from autogeo.utils.gemini import call_gemini | |
| from autogeo.utils.openai import call_openai | |
| from autogeo.utils.anthropic import call_claude | |
| from config_ui import get_missing_key_message | |
| def _detect_engine(engine_llm: str) -> str: | |
| """Map an engine_llm keyword to a canonical engine name.""" | |
| llm = engine_llm.lower() | |
| if "doubao" in llm or llm.startswith("ep-"): | |
| return "doubao" | |
| if "deepseek" in llm: | |
| return "deepseek" | |
| if "gpt" in llm: | |
| return "gpt" | |
| if "claude" in llm: | |
| return "claude" | |
| if "gemini" in llm: | |
| return "gemini" | |
| return "gemini" | |
| # --------------------------------------------------------------------------- | |
| # Public API | |
| # --------------------------------------------------------------------------- | |
| def run_rewrite( | |
| document: str, | |
| engine_llm: str = "doubao", | |
| dataset: str = "E-commerce", | |
| ) -> str: | |
| """Rewrite a document using the chosen engine.""" | |
| engine = _detect_engine(engine_llm) | |
| msg = get_missing_key_message(engine) | |
| if msg: | |
| raise ValueError(msg) | |
| return rewrite_document( | |
| document=document, | |
| dataset=dataset, | |
| engine_llm=engine_llm, | |
| ) | |
| def run_evaluate( | |
| query: str, | |
| text_list: List[str], | |
| target_id: int, | |
| engine_llm: str = "doubao", | |
| ) -> Dict: | |
| """Generate an LLM answer from *text_list* and compute GEO scores. | |
| Args: | |
| query: User question. | |
| text_list: List of 5 source documents (target at position *target_id*). | |
| target_id: Index of the document we care about. | |
| engine_llm: Model keyword. | |
| Returns: | |
| {"pos": float, "word": float, "wordpos": float} — GEO scores for target doc. | |
| """ | |
| engine = _detect_engine(engine_llm) | |
| # Call the right generate_answer_* function | |
| if engine == "deepseek": | |
| answer = generate_answer_deepseek(query, text_list, model_name=engine_llm) | |
| elif engine == "doubao": | |
| answer = generate_answer_doubao(query, text_list, model_name=engine_llm) | |
| elif engine == "gpt": | |
| answer = generate_answer_gpt(query, text_list, model_name=engine_llm) | |
| elif engine == "claude": | |
| answer = generate_answer_claude(query, text_list, model_name=engine_llm) | |
| else: | |
| answer = generate_answer_gemini(query, text_list, model_name="gemini-2.5-flash-lite") | |
| citations = extract_citations_new(answer) | |
| return { | |
| "pos": impression_pos_count_simple(citations)[target_id], | |
| "word": impression_word_count_simple(citations)[target_id], | |
| "wordpos": impression_wordpos_count_simple(citations)[target_id], | |
| } | |
| # --------------------------------------------------------------------------- | |
| # Iterative refinement — multi-round feedback loop | |
| # --------------------------------------------------------------------------- | |
| def _call_llm( | |
| user_prompt: str, | |
| engine_llm: str, | |
| system_prompt: str = "", | |
| temperature: float = 0.7, | |
| ) -> str: | |
| """Route a prompt to the right LLM client.""" | |
| engine = _detect_engine(engine_llm) | |
| msg = get_missing_key_message(engine) | |
| if msg: | |
| raise ValueError(msg) | |
| if engine == "deepseek": | |
| return call_deepseek(user_prompt, model_name=engine_llm, system_prompt=system_prompt, temperature=temperature) | |
| elif engine == "doubao": | |
| return call_doubao(user_prompt, model_name=engine_llm, system_prompt=system_prompt, temperature=temperature) | |
| elif engine == "gemini": | |
| return call_gemini(user_prompt, model_name="gemini-2.5-flash-lite", system_prompt=system_prompt, temperature=temperature) | |
| elif engine == "gpt": | |
| return call_openai(user_prompt, model_name=engine_llm, system_prompt=system_prompt, temperature=temperature) | |
| elif engine == "claude": | |
| return call_claude(user_prompt, model_name=engine_llm, system_prompt=system_prompt, temperature=temperature) | |
| else: | |
| # Fallback to Doubao | |
| return call_doubao(user_prompt, model_name="ep-20260525112247-ccwhk", system_prompt=system_prompt, temperature=temperature) | |
| def refine_document( | |
| original_doc: str, | |
| previous_rewrite: str, | |
| scores: Dict[str, float], | |
| query: str, | |
| engine_llm: str = "doubao", | |
| ) -> str: | |
| """Refine a rewrite based on evaluation feedback. | |
| Args: | |
| original_doc: The original document before any rewriting. | |
| previous_rewrite: The most recent rewritten version. | |
| scores: GEO scores from evaluating *previous_rewrite*. | |
| query: The user query for context. | |
| engine_llm: Model keyword. | |
| Returns: | |
| Improved rewritten document. | |
| """ | |
| # Build targeted feedback | |
| issues = [] | |
| if scores["pos"] < 0.3: | |
| issues.append("引用位置得分极低,开头部分需要立即抓住AI注意力,把核心信息放在最前面") | |
| elif scores["pos"] < 0.5: | |
| issues.append("引用位置得分偏低,考虑在开头直接给出结论性信息") | |
| if scores["word"] < 0.3: | |
| issues.append("引用篇幅严重不足,内容需要大幅扩展,增加具体数据、案例、细节") | |
| elif scores["word"] < 0.5: | |
| issues.append("引用篇幅偏少,可适当补充更多实质性内容") | |
| if scores["wordpos"] < 0.3: | |
| issues.append("综合引用得分很低,需要同时改善信息位置和内容密度") | |
| elif scores["wordpos"] < 0.5: | |
| issues.append("综合引用得分有提升空间") | |
| issues_text = "\n".join(f"- {i}" for i in issues) if issues else "各维度表现尚可,但仍有优化空间" | |
| system_prompt = ( | |
| "你是一位专业的 GEO(Generative Engine Optimization)优化专家。" | |
| "你的任务是根据评估反馈改进文档,使其在 AI 搜索引擎生成的答案中获得更高的引用可见性。" | |
| ) | |
| user_prompt = f"""请根据以下评估反馈,优化改写文档。 | |
| 【用户查询】 | |
| {query} | |
| 【原始文档】 | |
| {original_doc} | |
| 【当前改写版本】 | |
| {previous_rewrite} | |
| 【GEO 评估分数(满分 1.0)】 | |
| - 引用位置得分 (pos): {scores['pos']:.4f} — 衡量文档在AI答案中被引用的早晚 | |
| - 引用篇幅得分 (word): {scores['word']:.4f} — 衡量文档内容被引用的篇幅占比 | |
| - 综合引用得分 (wordpos): {scores['wordpos']:.4f} — 综合位置和篇幅的加权得分 | |
| 【需要改善的问题】 | |
| {issues_text} | |
| 请输出优化后的完整文档,保持中文输出,直接输出改写后的文档内容。""" | |
| return _call_llm( | |
| user_prompt=user_prompt, | |
| engine_llm=engine_llm, | |
| system_prompt=system_prompt, | |
| temperature=0.7, | |
| ) | |
| def run_rewrite_iterative( | |
| document: str, | |
| engine_llm: str, | |
| query: str, | |
| text_list_original: List[str], | |
| target_id: int, | |
| max_rounds: int = 3, | |
| progress_callback=None, | |
| ) -> List[Dict]: | |
| """Multi-round rewrite + evaluate with feedback. | |
| Round 1 uses the original CMU rewrite_document(). | |
| Rounds 2+ use refine_document() with evaluation feedback. | |
| Args: | |
| document: Original document text. | |
| engine_llm: Model keyword. | |
| query: User query for context. | |
| text_list_original: Full text list (target + distractors) for evaluation. | |
| target_id: Index of the target document. | |
| max_rounds: Maximum number of refinement rounds (1-5). | |
| progress_callback: Optional fn(round, status) for UI updates. | |
| Returns: | |
| List of {round, text, scores} dicts, one per round. | |
| """ | |
| results = [] | |
| # ---- Round 1: Standard rewrite (CMU original) ---- | |
| if progress_callback: | |
| progress_callback(1, "rewriting") | |
| rewritten = run_rewrite(document=document, engine_llm=engine_llm) | |
| if progress_callback: | |
| progress_callback(1, "evaluating") | |
| text_list_eval = [rewritten] + text_list_original[1:] | |
| scores = run_evaluate( | |
| query=query, | |
| text_list=text_list_eval, | |
| target_id=target_id, | |
| engine_llm=engine_llm, | |
| ) | |
| results.append({"round": 1, "text": rewritten, "scores": scores}) | |
| # ---- Rounds 2+: Refine with feedback ---- | |
| for rnd in range(2, max_rounds + 1): | |
| prev = results[-1] | |
| # Check if already excellent (all scores >= 0.9) | |
| if all(v >= 0.9 for v in prev["scores"].values()): | |
| break | |
| if progress_callback: | |
| progress_callback(rnd, "refining") | |
| refined = refine_document( | |
| original_doc=document, | |
| previous_rewrite=prev["text"], | |
| scores=prev["scores"], | |
| query=query, | |
| engine_llm=engine_llm, | |
| ) | |
| if progress_callback: | |
| progress_callback(rnd, "evaluating") | |
| text_list_eval = [refined] + text_list_original[1:] | |
| scores = run_evaluate( | |
| query=query, | |
| text_list=text_list_eval, | |
| target_id=target_id, | |
| engine_llm=engine_llm, | |
| ) | |
| results.append({"round": rnd, "text": refined, "scores": scores}) | |
| # Early stop if no improvement for 2 consecutive rounds | |
| if rnd >= 3: | |
| prev_prev = results[-2] | |
| prev_cur = results[-1] | |
| if prev_cur["scores"]["wordpos"] <= prev_prev["scores"]["wordpos"] * 1.02: | |
| break | |
| return results | |
| def generate_distractors( | |
| query: str, | |
| n: int = 4, | |
| engine_llm: str = "doubao", | |
| ) -> List[str]: | |
| """Use the selected LLM to generate competing search-result summaries. | |
| Args: | |
| query: The user query. | |
| n: Number of distractors to generate. | |
| engine_llm: Model keyword used for generation. | |
| Returns: | |
| List of *n* distractor strings. | |
| """ | |
| from openai import OpenAI | |
| from dotenv import load_dotenv | |
| load_dotenv(os.path.join(_parent, "keys.env")) | |
| engine = _detect_engine(engine_llm) | |
| msg = get_missing_key_message(engine) | |
| if msg: | |
| raise ValueError(msg) | |
| prompt = ( | |
| f"针对以下查询,生成 {n} 个不同的网站搜索结果摘要。" | |
| f"每个摘要从不同的角度回答查询,模拟真实搜索引擎中竞争网站的内容。" | |
| f"每个摘要 200-400 字,使用中文。\n\n" | |
| f"查询:{query}\n\n" | |
| f'请严格按照 JSON 数组格式输出,每个元素是一个摘要字符串。示例:["摘要1", "摘要2"]' | |
| ) | |
| if engine == "deepseek": | |
| client = OpenAI( | |
| api_key=os.getenv("DEEPSEEK_API_KEY"), | |
| base_url="https://api.deepseek.com/v1", | |
| ) | |
| model = engine_llm | |
| elif engine == "doubao": | |
| client = OpenAI( | |
| api_key=os.getenv("DOUBAO_API_KEY"), | |
| base_url="https://ark.cn-beijing.volces.com/api/v3", | |
| ) | |
| model = "ep-20260525112247-ccwhk" if not engine_llm.startswith("ep-") else engine_llm | |
| elif engine == "gpt": | |
| client = OpenAI(api_key=os.getenv("OPENAI_API_KEY")) | |
| model = engine_llm | |
| else: | |
| # Fallback: use DeepSeek as default distract generator | |
| client = OpenAI( | |
| api_key=os.getenv("DEEPSEEK_API_KEY"), | |
| base_url="https://api.deepseek.com/v1", | |
| ) | |
| model = "deepseek-chat" | |
| import json as _json | |
| resp = client.chat.completions.create( | |
| model=model, | |
| messages=[{"role": "user", "content": prompt}], | |
| temperature=0.8, | |
| ) | |
| raw = resp.choices[0].message.content.strip() | |
| # Strip markdown code fences if present (DeepSeek often wraps JSON in ```json ... ```) | |
| if raw.startswith("```"): | |
| raw = raw.split("\n", 1)[-1] if "\n" in raw else raw[3:] | |
| if raw.endswith("```"): | |
| raw = raw[:-3].strip() | |
| # Try to parse JSON array; if that fails, fall back to line splitting | |
| try: | |
| items = _json.loads(raw) | |
| if isinstance(items, list): | |
| return items[:n] | |
| except _json.JSONDecodeError: | |
| pass | |
| # Crude fallback: split by numbered items or double newlines | |
| import re | |
| parts = re.split(r"\n\d+[\.\、\)]\s*", raw) | |
| parts = [p.strip() for p in parts if len(p.strip()) > 20] | |
| return parts[:n] if parts else [raw[:400]] | |