Files
abap-llm/harness/trainset.py

158 lines
7.2 KiB
Python

"""Training mode of the task generator (stage 2 data).
Separate pool `tasks_gen/train` (ids G1000+). Same category and object type mix as the eval plan, plus
tasks for the common errors (docs/yol-haritasi.md step 3). A new bundle is checked against all eval tasks
and the earlier training tasks (`harness/overlap.py`); a too close bundle goes back to the model as a
repair message. K variants are not made here (they need the generic_v0 tool schema).
python3 -m harness.trainset plan
python3 -m harness.trainset run --part 0 --parts 3 [--target 200] [--stop-ledger 27]
"""
import argparse
import glob
import json
import os
import time
from .adt_client import load_env
from .evalset import SLOTS, RELEASES, accepted_goals
from .generator import ROOT, generate
from .ledger import BudgetExceeded, spent
from . import overlap
POOL = os.path.join(ROOT, "tasks_gen", "train")
PLAN = os.path.join(POOL, "plan.json")
FIRST_ID = 1000
RUN_BASE = 32000 # run numbers stay below 36**3 * 1 (proxy prefix rule); 40 per slot
N_SLOTS = 224 # about 12 % more slots than the target: some tasks are not accepted
N_ERROR = 30
EVAL_HOLD = 3 # eval pool is loaded once per process
ERROR_HINTS = [
("named-type", "The contract has methods with short coded character parameters (zone code, currency code, "
"status flag, region key) of fixed length. A correct solution declares named types for them: "
"TYPE c LENGTH n is not allowed in a method signature."),
("reserved-word", "The business domain uses concepts that collide with reserved or SQL words (hours, mode, "
"order, date, count, group, value). The contract uses safe field and parameter names (for example "
"work_hours, calc_mode); the solution must also avoid reserved words as field names."),
("long-names", "The task has at least 10 separate business rules, so a good solution needs at least 10 "
"own unit tests, each for one rule. Every ABAP name stays at 30 characters or less (the model must "
"shorten long descriptive test method names)."),
]
ERROR_TYPES = ["CLAS", "CLAS", "FUNC", "DDLS", "PROG", "CLAS"] # reserved-word fits DDLS and tables well
CATEGORY_SHARE = {"A": 10, "B": 15, "C": 15, "D": 10, "E": 15, "F": 10, "G": 10, "I": 5, "H": 10} # eval 8.1
ERROR_CATEGORY = {"named-type": "A", "reserved-word": "D", "long-names": "C"}
def build_plan():
"""Category shares as in the eval plan; object types inside a category as in evalset.SLOTS."""
share_sum = sum(CATEGORY_SHARE.values())
items = []
for cat, share in CATEGORY_SHARE.items():
k_cat = round(N_SLOTS * share / share_sum)
k_err = sum(1 for k in ERROR_CATEGORY if ERROR_CATEGORY[k] == cat) * N_ERROR // 3
slots = [x for x in SLOTS if x[0] == cat]
weight = sum(n for _, _, n, _ in slots)
k_rest = k_cat - k_err
quota = [(t, hints, round(k_rest * n / weight)) for _, t, n, hints in slots]
quota[0] = (quota[0][0], quota[0][1], quota[0][2] + k_rest - sum(q[2] for q in quota))
pos = 0
for t, hints, k in quota:
for i in range(k):
topic = hints[i % len(hints)] if hints else None
if cat == "G":
topic = (topic + "; " if topic else "") + f"release target {RELEASES[pos % 2]}"
pos += 1
items.append({"category": cat, "object_type": t, "topic": topic, "frac": pos / (k_cat + 1)})
for kind, text in ERROR_HINTS:
if ERROR_CATEGORY[kind] != cat:
continue
for i in range(N_ERROR // 3):
t = ERROR_TYPES[i % len(ERROR_TYPES)]
pos += 1
items.append({"category": cat, "object_type": t, "topic": text, "error_kind": kind,
"frac": pos / (k_cat + 1)})
items.sort(key=lambda x: (x["frac"], x["category"], x["object_type"])) # interleave: a cut keeps the mix
for n, it in enumerate(items):
it.pop("frac")
it["id"] = f"G{FIRST_ID + n:04d}"
it["run_base"] = RUN_BASE + 40 * n
return items
def accepted_count():
n = 0
for f in glob.glob(os.path.join(POOL, "G*", "generation.json")):
try:
n += bool(json.load(open(f)).get("accepted"))
except (OSError, ValueError):
pass
return n
def run(part, parts, target, stop_ledger):
os.makedirs(POOL, exist_ok=True)
if not os.path.exists(PLAN):
json.dump({"slots": build_plan(), "ledger_at_start": spent(), "stop_ledger": stop_ledger,
"target": target, "created": time.strftime("%F %T")}, open(PLAN, "w"), indent=1)
plan = json.load(open(PLAN))
stop_at = plan["ledger_at_start"] + plan["stop_ledger"]
base_url = os.environ.get("LLM_BASE_URL", "http://127.0.0.1:11434/v1")
evals = overlap.load_pool("eval")
for idx, s in enumerate(plan["slots"]):
if idx % parts != part:
continue
if os.path.exists(os.path.join(POOL, s["id"], "generation.json")):
continue
if accepted_count() >= plan["target"]:
print("TARGET reached", flush=True)
break
if spent() >= stop_at:
print("PHASE LIMIT", spent(), stop_at, flush=True)
break
goals = accepted_goals()
avoid = [g for g in goals if g] # eval goals and all training goals
topic = ((s["topic"] + ". ") if s["topic"] else "Choose a new, realistic business topic. ") + \
"Do not repeat these existing topics: " + "; ".join(avoid[-170:])
pool_now = evals + overlap.load_pool("train")
def extra(b, _pool=pool_now):
hits = overlap.check(overlap.load_bundle(b), _pool)
return [f"Too close to task {i} (similarity spec {sc['spec']:.2f}, rules {sc['core']:.2f}, "
f"names {sc['name']:.2f}). Choose a different business topic and different object names."
for i, sc in hits[:3]]
try:
log = generate(s["id"], "train", s["object_type"], s["category"], 2, "deepseek-v4.1-flash:cloud",
base_url, s["run_base"], topic, extra_check=extra)
except BudgetExceeded as e:
print("BUDGET", e, flush=True)
break
except Exception as e: # noqa: BLE001
log = {"id": s["id"], "error": str(e)[:500]}
log.update(error_kind=s.get("error_kind"), spent_total=spent())
print(json.dumps(log), flush=True)
def main():
load_env(os.path.join(ROOT, ".env"))
ap = argparse.ArgumentParser()
ap.add_argument("cmd", choices=["plan", "run"])
ap.add_argument("--part", type=int, default=0)
ap.add_argument("--parts", type=int, default=1)
ap.add_argument("--target", type=int, default=200)
ap.add_argument("--stop-ledger", type=float, default=27.0, help="ledger USD for this phase (10 USD usage = 27)")
a = ap.parse_args()
if a.cmd == "plan":
import collections
p = build_plan()
print(len(p), "slots", collections.Counter(x["category"] for x in p))
print(collections.Counter(x["object_type"] for x in p), collections.Counter(x.get("error_kind") for x in p))
return
run(a.part, a.parts, a.target, a.stop_ledger)
if __name__ == "__main__":
main()