Bläddra i källkod

feat(query): generate full f1 f2 combinations

SamLee 2 veckor sedan
förälder
incheckning
2038e6d0bd

+ 40 - 24
acquisition/queries/builder.py

@@ -4,6 +4,7 @@ from __future__ import annotations
 import json
 import random
 from dataclasses import dataclass
+from itertools import product
 from pathlib import Path
 from typing import Any
 
@@ -22,8 +23,8 @@ DEFAULT_ACTIVE_FAMILY_KEYS = ("f1", "f2")
 
 @dataclass(frozen=True)
 class QueryBuildOptions:
-    per: int = 30
-    batch_n: int = 30
+    per: int = 0
+    batch_n: int = 0
     seed: int = 7
     dry: bool = False
     active_family_keys: tuple[str, ...] = DEFAULT_ACTIVE_FAMILY_KEYS
@@ -52,7 +53,7 @@ def _leaves(idx: list[dict[str, Any]], source_type: str, under: str | None = Non
     return out
 
 
-def _nonleaf(
+def _nodes_at_depth(
     idx: list[dict[str, Any]],
     source_type: str,
     depths: tuple[int, ...] = (3, 4),
@@ -61,21 +62,24 @@ def _nonleaf(
     paths = [(_segs(n.get("path")), n.get("name")) for n in idx if n.get("source_type") == source_type]
     if under:
         paths = [(s, nm) for s, nm in paths if under in s]
-    all_paths = {"/".join(s) for s, _ in paths}
     out: list[str] = []
     seen: set[str] = set()
     for segs, name in paths:
         if len(segs) not in depths:
             continue
-        full = "/".join(segs)
-        is_nonleaf = any(other != full and other.startswith(full + "/") for other in all_paths)
         value = name or (segs[-1] if segs else "")
-        if is_nonleaf and value and value not in seen:
+        if value and value not in seen:
             seen.add(value)
             out.append(value)
     return out
 
 
+def _sample_axis(values: list[str], limit: int, rng: random.Random) -> list[str]:
+    if limit <= 0 or limit >= len(values):
+        return values
+    return rng.sample(values, limit)
+
+
 def build_creation_query_batch(
     settings: Settings,
     *,
@@ -86,20 +90,34 @@ def build_creation_query_batch(
     opts = options or QueryBuildOptions()
     rng = random.Random(opts.seed)
     idx = json.loads(tree_path.read_text("utf-8"))
-    shi = _nonleaf(idx, "实质", depths=(3, 4), under="理念")
-    xing = _nonleaf(idx, "形式", depths=(3, 4), under="架构")
+    shi = _nodes_at_depth(idx, "实质", depths=(3, 4))
+    xing = _nodes_at_depth(idx, "形式", depths=(3, 4))
     purpose_pool = _leaves(idx, "作用") + _leaves(idx, "感受") + _leaves(idx, "意图")
-    shi_batch = rng.sample(shi, min(opts.batch_n, len(shi)))
-    xing_batch = rng.sample(xing, min(opts.batch_n, len(xing)))
-    purpose_batch = rng.sample(purpose_pool, min(opts.batch_n, len(purpose_pool)))
+    shi_batch = _sample_axis(shi, opts.batch_n, rng)
+    xing_batch = _sample_axis(xing, opts.batch_n, rng)
+    purpose_batch = _sample_axis(purpose_pool, opts.batch_n, rng)
     f6_zy = _nonleaf_d4("作用", 10, tree_path=tree_path)
     f6_stage_act = [(s, a) for s in STAGES for a in ACTIONS] + [("", "")]
 
     if not shi_batch or not xing_batch or not purpose_batch or not f6_zy:
         raise RuntimeError("scope tree does not contain enough creation query axes")
 
+    product_pools = {
+        "实质": shi_batch,
+        "形式": xing_batch,
+        "目的": purpose_batch,
+        "模态": MODALITY,
+        "业务阶段": INTENT,
+        "知识类型": KTYPE,
+    }
+
+    def gen_product(keys: list[str]):
+        for values in product(*(product_pools[key] for key in keys)):
+            yield {"parts": dict(zip(keys, values, strict=True))}
+
+    legacy_per = opts.per if opts.per > 0 else 30
     master: list[dict[str, str]] = []
-    for i in range(opts.per):
+    for i in range(legacy_per):
         stage, action = f6_stage_act[i % len(f6_stage_act)]
         master.append(
             {
@@ -117,10 +135,6 @@ def build_creation_query_batch(
             }
         )
 
-    def project(i: int, keys: list[str]) -> dict[str, Any]:
-        row = master[i]
-        return {"parts": {key: row[key] for key in keys}}
-
     def gen_old(
         i: int,
         *,
@@ -152,11 +166,11 @@ def build_creation_query_batch(
         return {"parts": parts, "query": query}
 
     families = [
-        {"key": "f1", "axes": ["实质", "模态", "业务阶段", "知识类型"], "gen": lambda i: project(i, ["实质", "模态", "业务阶段", "知识类型"])},
-        {"key": "f2", "axes": ["形式", "模态", "业务阶段", "知识类型"], "gen": lambda i: project(i, ["形式", "模态", "业务阶段", "知识类型"])},
-        {"key": "f4", "axes": ["作用/感受/意图", "模态", "业务阶段", "知识类型"], "gen": lambda i: project(i, ["目的", "模态", "业务阶段", "知识类型"])},
-        {"key": "f3", "axes": ["实质", "形式", "模态", "业务阶段", "知识类型"], "gen": lambda i: project(i, ["实质", "形式", "模态", "业务阶段", "知识类型"])},
-        {"key": "f5", "axes": ["模态", "业务阶段", "知识类型"], "gen": lambda i: project(i, ["模态", "业务阶段", "知识类型"])},
+        {"key": "f1", "axes": ["实质", "模态", "业务阶段", "知识类型"], "items": lambda: gen_product(["实质", "模态", "业务阶段", "知识类型"])},
+        {"key": "f2", "axes": ["形式", "模态", "业务阶段", "知识类型"], "items": lambda: gen_product(["形式", "模态", "业务阶段", "知识类型"])},
+        {"key": "f4", "axes": ["作用/感受/意图", "模态", "业务阶段", "知识类型"], "items": lambda: gen_product(["目的", "模态", "业务阶段", "知识类型"])},
+        {"key": "f3", "axes": ["实质", "形式", "模态", "业务阶段", "知识类型"], "items": lambda: gen_product(["实质", "形式", "模态", "业务阶段", "知识类型"])},
+        {"key": "f5", "axes": ["模态", "业务阶段", "知识类型"], "items": lambda: gen_product(["模态", "业务阶段", "知识类型"])},
         {"key": "a_shi", "axes": ["实质", "阶段", "动作", "作用", "知识类型"], "gen": lambda i: gen_old(i, shi_axis=True)},
         {"key": "a_xing", "axes": ["形式", "阶段", "动作", "作用", "知识类型"], "gen": lambda i: gen_old(i, xing_axis=True)},
         {"key": "a_both", "axes": ["实质", "形式", "阶段", "动作", "作用", "知识类型"], "gen": lambda i: gen_old(i, shi_axis=True, xing_axis=True)},
@@ -202,14 +216,16 @@ def build_creation_query_batch(
         name = " × ".join(family["axes"])
         seen: set[str] = set()
         items: list[dict[str, Any]] = []
-        for i in range(opts.per):
-            generated = family["gen"](i)
+        generated_items = family["items"]() if "items" in family else (family["gen"](i) for i in range(legacy_per))
+        for generated in generated_items:
             parts = generated["parts"]
             query = generated.get("query") or " ".join(parts[k] for k in order if k in parts)
             if query in seen:
                 continue
             seen.add(query)
             items.append({"query": query, "parts": parts})
+            if opts.per > 0 and len(items) >= opts.per:
+                break
         verdicts = (
             [
                 {"keep": True, "valid": None, "relevant": True, "reason": "(dry:未过筛)"}

+ 7 - 2
acquisition/queries/filter.py

@@ -9,6 +9,11 @@ from pathlib import Path
 import httpx
 
 from core.config import Settings, load_env_file
+from core.text_limits import (
+    ERROR_MESSAGE_MAX_CHARS,
+    QUERY_FILTER_REASON_MAX_CHARS,
+    clip_text,
+)
 
 PROMPT = Path(__file__).resolve().parent.parent / "query_filter.txt"
 VALID_MIN = 6
@@ -65,12 +70,12 @@ def _filter_batch(queries: list[str], settings: Settings) -> list[dict]:
                     "keep": valid >= VALID_MIN and relevant,
                     "valid": valid,
                     "relevant": relevant,
-                    "reason": str(d.get("reason", ""))[:50],
+                    "reason": clip_text(d.get("reason", ""), QUERY_FILTER_REASON_MAX_CHARS),
                 }
             )
         return out
     except Exception as exc:
-        reason = f"筛选失败:{str(exc)[:30]}"
+        reason = f"筛选失败:{clip_text(exc, ERROR_MESSAGE_MAX_CHARS)}"
         return [
             {"keep": True, "valid": 10, "relevant": True, "reason": reason}
             for _ in queries

+ 1 - 1
acquisition/query_filter.txt

@@ -25,4 +25,4 @@ E.【成品 / 作品 / 素材本身】一篇具体作品、一份素材 / 金句
 (是否最终保留由调用方按"valid ≥ 阈值 且 relevant=true"决定,你只需如实打分 / 判断。)
 
 严格只输出 JSON 数组,idx 对应输入序号,无任何额外文字 / 解释 / markdown 围栏:
-[{"idx":0,"valid":8,"relevant":true,"reason":"≤30字,说明打分与判断依据"}]
+[{"idx":0,"valid":8,"relevant":true,"reason":"说明打分与判断依据;保留必要证据,不要为了简短省略关键原因"}]

+ 2 - 2
scripts/build_creation_demo.py

@@ -30,8 +30,8 @@ def parse_args(argv: list[str] | None = None) -> argparse.Namespace:
     parser = argparse.ArgumentParser(description=__doc__)
     parser.add_argument("--env-file", default=".env")
     parser.add_argument("--tree-path", type=Path)
-    parser.add_argument("--per", type=int, default=30)
-    parser.add_argument("--batch-n", type=int, default=30)
+    parser.add_argument("--per", type=int, default=0, help="Max queries per family; 0 means all combinations")
+    parser.add_argument("--batch-n", type=int, default=0, help="Max tree nodes per axis; 0 means all L3/L4 nodes")
     parser.add_argument("--seed", type=int, default=7)
     parser.add_argument("--dry", action="store_true", help="Skip LLM query filtering")
     parser.add_argument(

+ 68 - 0
tests/test_query_builder.py

@@ -1,16 +1,20 @@
 from __future__ import annotations
 
+import json
 from uuid import uuid4
 
 import pytest
 
 from acquisition.domain import Query, QueryBatch
+from acquisition.queries import filter as query_filter
 from acquisition.queries.builder import (
     QueryBuildOptions,
+    TREES,
     build_creation_query_batch,
     persist_query_batch,
 )
 from core.config import PgConfig, Settings
+from core.text_limits import QUERY_FILTER_REASON_MAX_CHARS
 
 
 def _settings() -> Settings:
@@ -30,6 +34,23 @@ def _settings() -> Settings:
     )
 
 
+def test_query_filter_preserves_wide_reason(monkeypatch):
+    reason = "理" * (QUERY_FILTER_REASON_MAX_CHARS + 3)
+
+    monkeypatch.setattr(
+        query_filter,
+        "_chat_content",
+        lambda settings, messages: json.dumps(
+            [{"idx": 0, "valid": 8, "relevant": True, "reason": reason}],
+            ensure_ascii=False,
+        ),
+    )
+
+    rows = query_filter.filter_queries(["符号 视频 灵感 怎么做"], _settings())
+
+    assert rows[0]["reason"] == "理" * QUERY_FILTER_REASON_MAX_CHARS
+
+
 class FakeRepo:
     def __init__(self) -> None:
         self.batch_kwargs = None
@@ -44,6 +65,18 @@ class FakeRepo:
         return Query(id=uuid4(), **kwargs)
 
 
+def _expected_l3_l4_names(source_type: str) -> set[str]:
+    rows = json.loads(TREES.read_text("utf-8"))
+    names: set[str] = set()
+    for row in rows:
+        if row.get("source_type") != source_type:
+            continue
+        depth = len([part for part in (row.get("path") or "").split("/") if part])
+        if depth in (3, 4) and row.get("name"):
+            names.add(row["name"])
+    return names
+
+
 def test_build_creation_query_batch_defaults_to_first_two_families():
     generated = build_creation_query_batch(
         _settings(),
@@ -55,6 +88,41 @@ def test_build_creation_query_batch_defaults_to_first_two_families():
     assert sum(len(family["items"]) for family in generated["families"]) == 4
 
 
+def test_build_creation_query_batch_uses_all_l3_l4_substance_and_form_nodes():
+    generated = build_creation_query_batch(
+        _settings(),
+        options=QueryBuildOptions(per=2, batch_n=999, dry=True),
+    )
+
+    assert set(generated["axis_values"]["实质"]) == _expected_l3_l4_names("实质")
+    assert set(generated["axis_values"]["形式"]) == _expected_l3_l4_names("形式")
+    assert {"原声录音", "动物声音", "场景原声", "制作音效", "环境音", "符号", "文字符号"} <= set(
+        generated["axis_values"]["实质"]
+    )
+    assert {"配乐", "语音"} <= set(generated["axis_values"]["形式"])
+
+
+def test_build_creation_query_batch_expands_active_families_as_cartesian_products():
+    generated = build_creation_query_batch(
+        _settings(),
+        options=QueryBuildOptions(per=0, batch_n=0, dry=True),
+    )
+
+    by_key = {family["key"]: family for family in generated["families"]}
+    f1_items = by_key["f1"]["items"]
+    f2_items = by_key["f2"]["items"]
+
+    assert len(f1_items) == len(generated["axis_values"]["实质"]) * 2 * 3 * 3
+    assert len(f2_items) == len(generated["axis_values"]["形式"]) * 2 * 3 * 3
+    assert {
+        item["parts"]["知识类型"]
+        for item in f1_items
+        if item["parts"]["实质"] == "公共安全"
+        and item["parts"]["模态"] == "视频"
+        and item["parts"]["业务阶段"] == "灵感"
+    } == {"怎么做", "有哪些", "为什么"}
+
+
 def test_build_creation_query_batch_can_explicitly_enable_reserved_families():
     all_keys = (
         "f1",