laya-browser / code /finetune /convert_webchain.py
cklxx's picture
v19s: WebChain real-site trajectories, format v5, webgym x7 + DAgger, harness fixes; replaces v17s
454b3e6 verified
Raw History Blame Contribute Delete
15.6 kB
"""WebChain (webagentlab/webchain, CC-BY-4.0; human trajectories on real sites) -> laya cases.
python finetune/convert_webchain.py fetch <n_traces> [workers=16] # AX-tree snapshots -> compact .json.gz (direct, no proxy)
python finetune/convert_webchain.py convert <out.jsonl> # compact snapshots -> cases (same format as m2w_cases)
Per step: goal = the trace's query; page = the AX snapshot taken at that step (interactive nodes + text-bearing nodes,
page order); gold = the node the human acted on, found by its text (`value`) + html tag, ties broken by the selector's
id/class tokens -- an ambiguous or missing gold drops the step. Candidates mirror Mind2Web's: the interactive nodes plus
some non-semantic text nodes (the gold is often a clickable <span>/<div>), capped at 60 in page order.
Actions kept: click / double_click -> CLICK, type -> TYPE_TEXT (or SELECT on a native <select>), press_enter -> PRESS_ENTER.
hover / drag / copy / paste / right_click steps are dropped (no laya operation), but still count in the history.
WebChain records no scroll and no final DONE, so it contributes no DONE labels.
"""
import gzip, hashlib, json, os, random, re, sys, time
from concurrent.futures import ThreadPoolExecutor
META = os.path.expanduser("~/data/webchain/data/seed_sft/metadata")
SNAP = os.path.expanduser("~/data/webchain/ax_compact")
INTERACTIVE = {"link", "button", "textbox", "searchbox", "combobox", "checkbox", "radio", "switch", "tab", "menuitem",
"menuitemradio", "menuitemcheckbox", "option", "listbox", "slider", "spinbutton", "treeitem", "gridcell"}
TEXTY = {"generic", "listitem", "cell", "heading", "img", "paragraph", "StaticText", "label", "row", "columnheader"}
KEEP_ATTR = ("data-imean-axt-id", "html_tag", "href", "type", "placeholder", "aria-label", "title", "alt", "value", "id", "class", "aria-checked",
"aria-selected", "aria-expanded", "checked", "selected", "role")
def compact(ax):
"""Flatten an AX snapshot to [{role,name,a,vis,d}] in document order (iterative: real pages nest deeper than
Python's recursion limit)."""
out, stack = [], [(ax, 0)]
while stack:
n, depth = stack.pop()
a = n.get("attributes") or {}
vis = (n.get("offsetWidth") or 0) > 0 and (n.get("offsetHeight") or 0) > 0
out.append({"role": n.get("role") or "", "name": (n.get("name") or "")[:200], "d": depth, "vis": vis,
"a": {k: str(a[k])[:120] for k in KEEP_ATTR if k in a}})
stack.extend((c, depth + 1) for c in reversed(n.get("children") or []))
return out
def gold_axt_from_dom(html, selector):
"""The data-imean-axt-id of the element the CSS selector picks (or its nearest tagged ancestor / first tagged child)."""
import lxml.html
try:
root = lxml.html.fromstring(html)
els = root.cssselect(selector)
except Exception:
return None
if len(els) != 1:
return None
el = els[0]
for e in [el, *el.iterancestors()]:
if e.get("data-imean-axt-id"):
return e.get("data-imean-axt-id")
for e in el.iterdescendants():
if e.get("data-imean-axt-id"):
return e.get("data-imean-axt-id")
return None
def step_file(uid, idx):
return os.path.join(SNAP, hashlib.md5(f"{uid}:{idx}".encode()).hexdigest() + ".json.gz")
def fetch(n_traces, workers):
"""Per step: the AX snapshot (compact) + the gold node's axt id. The gold is found by text in the AX tree; only when
that fails is the (bigger) DOM snapshot fetched and the step's CSS selector resolved to an axt id."""
import pandas as pd, httpx
os.makedirs(SNAP, exist_ok=True)
tr = pd.read_parquet(f"{META}/traces.parquet", columns=["uid"])
uids = sorted(tr["uid"].tolist()); random.Random(0).shuffle(uids)
keep = set(uids[:n_traces])
ac = pd.read_parquet(f"{META}/actions.parquet", columns=["trace_uid", "source_step_index", "action_type", "value", "title", "attributes",
"selector", "ax_tree_url", "html_dom_url"])
rows = [r for r in ac[ac.trace_uid.isin(keep)].to_dict("records") if r["action_type"] in ("click", "double_click", "type", "select", "press_enter")
and isinstance(r["ax_tree_url"], str) and r["ax_tree_url"].startswith("http")]
client = httpx.Client(timeout=60, trust_env=False) # direct: never through the proxy
stats = {"ok": 0, "skip": 0, "err": 0, "dom": 0, "gold_text": 0, "gold_dom": 0, "no_gold": 0, "mb": 0.0}
def get(u):
for attempt in range(3):
try:
r = client.get(u); r.raise_for_status(); stats["mb"] += len(r.content) / 1e6; return r
except Exception:
time.sleep(2 * (attempt + 1))
return None
def one(row):
dst = step_file(row["trace_uid"], row["source_step_index"])
if os.path.exists(dst):
stats["skip"] += 1; return
try:
r = get(row["ax_tree_url"])
if r is None:
stats["err"] += 1; return
nodes = compact(r.json()); gold_axt = None
if row["action_type"] != "press_enter":
# DOM selector first (exact); text matching only when the selector does not resolve (it disagrees with
# the selector on ~13% of steps, so it is the fallback, not the default)
if isinstance(row["html_dom_url"], str) and row["html_dom_url"].startswith("http") and row["selector"]:
d = get(row["html_dom_url"]); stats["dom"] += 1
gold_axt = gold_axt_from_dom(d.content, row["selector"]) if d is not None else None
if gold_axt and not any(n["a"].get("data-imean-axt-id") == gold_axt for n in nodes):
gold_axt = None
if gold_axt:
stats["gold_dom"] += 1
if not gold_axt:
g = find_gold(nodes, row)
if g is not None:
gold_axt = nodes[g]["a"].get("data-imean-axt-id"); stats["gold_text"] += 1
else:
stats["no_gold"] += 1
with gzip.open(dst + ".tmp", "wt") as f:
json.dump({"nodes": nodes, "gold_axt": gold_axt}, f, ensure_ascii=False, separators=(",", ":"))
os.replace(dst + ".tmp", dst); stats["ok"] += 1
except Exception as e:
stats["err"] += 1
if os.path.exists(dst + ".tmp"):
os.remove(dst + ".tmp")
t0 = time.time()
with ThreadPoolExecutor(workers) as ex:
for i, _ in enumerate(ex.map(one, rows)):
if i % 1000 == 0:
print(f" {i}/{len(rows)} {stats} {stats['mb'] / max(1, time.time() - t0):.1f} MB/s", flush=True)
print("fetch done", len(keep), "traces,", len(rows), "steps", stats)
def _norm(s):
return " ".join(str(s or "").split()).lower()
def label_of(n):
a = n["a"]
return n["name"] or a.get("aria-label") or a.get("placeholder") or a.get("title") or a.get("alt") or a.get("value") or ""
def find_gold(nodes, row):
"""Index of the acted-on node, or None if missing / ambiguous."""
val = _norm(row["value"]); ttl = _norm(re.sub(r"^(点击|输入|选择|click|type|select)\s*", "", str(row["title"] or ""), flags=re.I))
try:
tag = (json.loads(row["attributes"] or "{}").get("data", {}).get("node", {}).get("name") or "").lower()
except Exception:
tag = ""
sel = str(row["selector"] or "")
toks = set(re.findall(r"[#.]([A-Za-z0-9_-]+)", sel))
typing = row["action_type"] in ("type", "select")
cand = []
for i, n in enumerate(nodes):
if typing:
if n["a"].get("html_tag") not in ("input", "textarea", "select") and n["role"] not in ("textbox", "searchbox", "combobox"):
continue
lab = _norm(label_of(n))
score = 3 * (tag and n["a"].get("html_tag") == tag) + 2 * bool(toks & ({n["a"].get("id", "")} | set(n["a"].get("class", "").split())))
score += 2 * bool(lab and (lab in ttl or ttl in lab))
cand.append((score, i))
else:
lab = _norm(label_of(n))
if not lab or not (lab == val or (ttl and lab == ttl)):
continue
score = 3 * (tag and n["a"].get("html_tag") == tag) + 2 * bool(toks & ({n["a"].get("id", "")} | set(n["a"].get("class", "").split())))
score += n["role"] in INTERACTIVE
cand.append((score, i))
if not cand:
return None
cand.sort(reverse=True)
if len(cand) > 1 and cand[0][0] == cand[1][0]:
return None
return cand[0][1]
def page_from(nodes, gold, rng, url, title):
"""(page_obj, gold action id, gold kind, gold label, options) in m2w_cases format."""
inter = [i for i, n in enumerate(nodes) if n["vis"] and n["role"] in INTERACTIVE and label_of(n)]
texty = [i for i, n in enumerate(nodes) if n["vis"] and n["role"] in TEXTY and 2 <= len(label_of(n)) <= 80 and i != gold]
pick = set(inter[:45]) | set(rng.sample(texty, min(len(texty), 15))) | {gold}
pick = sorted(pick)[:60] if gold in sorted(pick)[:60] else sorted(set(sorted(pick)[:59]) | {gold})
acts = []
for j, i in enumerate(pick):
n = nodes[i]; tag = n["a"].get("html_tag"); lab = label_of(n)[:120]
editable = n["role"] in ("textbox", "searchbox") or (n["role"] == "combobox" and tag in ("input", "textarea"))
base = {"node": j + 1, "label": lab, "role": n["role"] or tag or "generic"}
for k in ("aria-checked", "aria-selected", "aria-expanded"):
if k in n["a"]:
base[k[5:]] = n["a"][k]
if editable:
acts.append({**base, "id": f"fill:{j + 1}", "kind": "fill", "value": n["a"].get("value", ""), "current_value": n["a"].get("value", "")})
acts.append({**base, "id": f"click:{j + 1}", "kind": "click"})
gj = pick.index(gold) + 1
text, seen = [], set()
for n in nodes:
t = n["name"].strip()
if n["vis"] and t and n["role"] in TEXTY | {"link", "button"} and t not in seen and len(t) < 300:
seen.add(t); text.append(t)
return {"url": url, "title": title, "text": "\n".join(text)[:6000], "actions": acts}, gj
def convert(out_path):
import pandas as pd
tr = pd.read_parquet(f"{META}/traces.parquet", columns=["uid", "query", "primary_host", "intent_type", "web_type"])
q = {r.uid: r for r in tr.itertuples(index=False)}
ac = pd.read_parquet(f"{META}/actions.parquet", columns=["trace_uid", "source_step_index", "action_type", "input_text", "value", "title",
"attributes", "selector", "ax_tree_url", "href", "host_title"])
ac = ac.sort_values(["trace_uid", "source_step_index"])
stats = {"steps": 0, "no_snapshot": 0, "no_gold": 0, "skipped_op": 0, "cases": 0}
ops = {"click": "CLICK", "double_click": "CLICK", "type": "TYPE_TEXT", "select": "TYPE_TEXT", "press_enter": "PRESS_ENTER"}
with open(out_path, "w") as out:
for uid, g in ac.groupby("trace_uid", sort=False):
hist = []
for row in g.to_dict("records"):
stats["steps"] += 1
at = row["action_type"]; txt = row["input_text"] if row["input_text"] not in (None, "", "no input text") else None
snap = step_file(uid, row["source_step_index"])
done_hist = {"action": (str(row["value"] or row["title"] or at))[:80], "kind": {"type": "fill", "select": "select"}.get(at, "click"),
"text": txt, "page_changed": True}
if at not in ops:
stats["skipped_op"] += 1; hist.append(done_hist); continue
if not os.path.exists(snap):
stats["no_snapshot"] += 1; hist.append(done_hist); continue
blob = json.load(gzip.open(snap, "rt")); nodes = blob["nodes"]
rng = random.Random(hash((uid, row["source_step_index"])) & 0xffffffff)
url, title = str(row["href"] or ""), str(row["host_title"] or "")
t = q[uid]
base_case = {"page": -1, "url": url, "title": title, "goal": re.sub(r'^\s*(task\s*\d+\s*[::.]|\d+\s*[.、)])\s*', "", str(t.query).strip(), flags=re.I).strip().strip('"').strip(), "history": list(hist)[-10:], "source": "webchain",
"task_id": uid, "website": str(t.primary_host), "intent": str(t.intent_type), "domain": str(t.web_type)}
if at == "press_enter":
if not hist:
stats["no_gold"] += 1; hist.append(done_hist); continue
pg, _ = page_from(nodes, max(0, len(nodes) - 1), rng, url, title)
pg["actions"].append({"id": "press_enter", "kind": "key", "label": "Press Enter in the focused text field (submit it)", "key": "Enter"})
out.write(json.dumps({**base_case, "page_obj": pg, "gold_op": "PRESS_ENTER", "gold_id": "press_enter", "kind": "key", "label": ""}, ensure_ascii=False) + "\n")
stats["cases"] += 1; hist.append(done_hist); continue
gold = next((i for i, n in enumerate(nodes) if blob["gold_axt"] and n["a"].get("data-imean-axt-id") == blob["gold_axt"]), None)
if gold is None:
stats["no_gold"] += 1; hist.append(done_hist); continue
pg, gj = page_from(nodes, gold, rng, url, title)
gn = nodes[gold]
if at in ("type", "select"):
if gn["a"].get("html_tag") == "select":
op, gid, kind = "TYPE_TEXT", f"fill:{gj}", "fill" # native select options are not in the AX snapshot
else:
op, gid, kind = "TYPE_TEXT", f"fill:{gj}", "fill"
if not any(a["id"] == gid for a in pg["actions"]):
# the typed-into node is not an editable role in the snapshot: offer it as a text field
a0 = next(a for a in pg["actions"] if a["node"] == gj)
pg["actions"].insert(pg["actions"].index(a0), {**{k: v for k, v in a0.items() if k not in ("id", "kind")}, "id": gid, "kind": "fill", "value": "", "current_value": ""})
else:
op, gid, kind = "CLICK", f"click:{gj}", "click"
lab = next(a["label"] for a in pg["actions"] if a["id"] == gid)
# an unlabeled gold (icon button, svg) or one whose label repeats among the candidates cannot be learned
# from text: drop the step (35% / 13% of gold clicks before this filter)
labs = [a["label"].strip().lower() for a in pg["actions"] if a["kind"] == ("fill" if kind == "fill" else "click")]
if not lab.strip() or labs.count(lab.strip().lower()) > 1:
stats["unlearnable"] = stats.get("unlearnable", 0) + 1; hist.append({**done_hist, "action": lab[:80] or done_hist["action"]}); continue
out.write(json.dumps({**base_case, "page_obj": pg, "gold_op": op, "gold_id": gid, "kind": kind, "label": lab, "gold_text": txt}, ensure_ascii=False) + "\n")
stats["cases"] += 1
hist.append({**done_hist, "action": lab[:80]})
print("convert done", stats, "->", out_path)
if __name__ == "__main__":
if sys.argv[1] == "fetch":
fetch(int(sys.argv[2]), int(sys.argv[3]) if len(sys.argv) > 3 else 16)
else:
convert(sys.argv[2])