AutoGEO-Studio / app /engine.py
chaseurstep's picture
Upload folder using huggingface_hub
5a45328 verified
Raw
History Blame Contribute Delete
13.3 kB
"""
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]]