contimp-app / app /tasks /config_copilot.py
lefft's picture
Per-task LLM backends (OpenAI + Anthropic wire); config-copilot pinned to the oumi backend
1bcb9d8 verified
Raw
History Blame Contribute Delete
14.3 kB
"""Task 2: turn a natural-language training-run request into a valid oumi YAML.
The agent gets two tools: get_schema (discover the supported config subset) and
validate_config (YAML parse + JSON Schema check). Dealt requests are rendered
from a sampled field spec, so ground truth is the spec itself: scoring checks
schema validity plus exact per-field match on the targeted values.
"""
import json
import random
import re
import jsonschema
import yaml
from app.config import DATA_DIR
from app.tasks.base import Sample, Task
SCHEMA = json.loads((DATA_DIR / "oumi_training_schema.json").read_text())
_VALIDATOR = jsonschema.Draft202012Validator(SCHEMA)
SYSTEM_PROMPT = """\
You are the Oumi Config Copilot. You turn a user's plain-English description of a \
fine-tuning run into a valid oumi training config (YAML).
Rules:
- Only use fields that exist in the config schema. Call get_schema to see the \
schema for a section before using unfamiliar fields.
- Supervised fine-tuning uses training.trainer_type: TRL_SFT.
- LoRA/QLoRA runs set training.use_peft: true and configure the peft section. \
QLoRA additionally sets peft.q_lora: true.
- For a local JSONL file of chat conversations, use dataset_name: "text_sft_jsonl" \
with dataset_path pointing at the file.
- Include the fields the user asked for plus anything required by the schema. \
Don't pad the config with unrelated defaults.
- Before giving your final answer, call validate_config on your draft and fix any \
errors it reports.
- Your final message must be exactly one fenced ```yaml code block containing the \
config, with no text after it.\
"""
TOOLS = [
{
"type": "function",
"function": {
"name": "get_schema",
"description": "Return the JSON Schema for a section of the oumi training config (or the whole thing).",
"parameters": {
"type": "object",
"properties": {
"section": {
"type": "string",
"enum": ["all", "model", "data", "training", "peft"],
}
},
"required": ["section"],
},
},
},
{
"type": "function",
"function": {
"name": "validate_config",
"description": "Validate a candidate oumi training config. Returns {valid, errors}.",
"parameters": {
"type": "object",
"properties": {
"yaml_config": {"type": "string", "description": "The full YAML config text"}
},
"required": ["yaml_config"],
},
},
},
]
def validate(config: object) -> list[str]:
errors = sorted(_VALIDATOR.iter_errors(config), key=lambda e: list(e.absolute_path))
return [
f"at {'.'.join(str(p) for p in e.absolute_path) or '<root>'}: {e.message}"
for e in errors
][:12]
def execute_tool(name: str, args: dict) -> str:
if name == "get_schema":
section = args.get("section", "all")
sub = SCHEMA if section == "all" else SCHEMA["properties"].get(section, {})
return json.dumps(sub, indent=1)
if name == "validate_config":
try:
config = yaml.safe_load(args.get("yaml_config", ""))
except yaml.YAMLError as e:
return json.dumps({"valid": False, "errors": [f"YAML parse error: {e}"]})
if not isinstance(config, dict):
return json.dumps({"valid": False, "errors": ["config must be a YAML mapping"]})
errors = validate(config)
return json.dumps({"valid": not errors, "errors": errors})
return json.dumps({"error": f"unknown tool {name}"})
# --- Request generation ------------------------------------------------------
MODELS = [
"Qwen/Qwen3-4B-Instruct-2507", "Qwen/Qwen2.5-1.5B-Instruct",
"meta-llama/Llama-3.2-3B-Instruct", "meta-llama/Llama-3.1-8B-Instruct",
"HuggingFaceTB/SmolLM2-1.7B-Instruct", "google/gemma-3-4b-it",
"microsoft/Phi-4-mini-instruct", "mistralai/Mistral-7B-Instruct-v0.3",
]
HUB_DATASETS = [
"yahma/alpaca-cleaned", "HuggingFaceH4/ultrachat_200k",
"databricks/databricks-dolly-15k", "OpenAssistant/oasst1",
]
LOCAL_PATHS = [
"data/train.jsonl", "/mnt/data/conversations.jsonl", "datasets/support_chats.jsonl",
"data/curated/sft_round2.jsonl", "exports/traces_clean.jsonl",
]
PROJECT_SLUGS = ["support-bot", "triage-v2", "summarizer", "qa-assistant", "intent-router"]
OPENERS = [
"", "Hey! ", "Quick one: ", "Setting up a new experiment. ",
"Time for another ablation. ", "For the {slug} project: ",
"Need a config for tomorrow's run. ",
]
def generate(seed: int) -> tuple[dict, str]:
"""Deterministically generate (target field spec, natural-language request)."""
rng = random.Random(seed)
fields: dict[str, object] = {}
phrases: list[str] = []
fields["training.trainer_type"] = "TRL_SFT"
model = rng.choice(MODELS)
fields["model.model_name"] = model
kind = rng.choice(["full", "lora", "lora", "qlora"]) # lora-heavy mix
kind_phrase = {
"full": rng.choice([f"do a full-parameter SFT run on {model}",
f"fully fine-tune {model} (no adapters)"]),
"lora": rng.choice([f"fine-tune {model} with LoRA",
f"do a LoRA SFT run on {model}"]),
"qlora": rng.choice([f"fine-tune {model} with QLoRA (4-bit)",
f"do a 4-bit QLoRA run on {model}"]),
}[kind]
fields["training.use_peft"] = kind != "full"
if kind == "qlora":
fields["peft.q_lora"] = True
fields["peft.q_lora_bits"] = 4
if kind in ("lora", "qlora"):
if rng.random() < 0.7:
r = rng.choice([8, 16, 32, 64])
fields["peft.lora_r"] = r
phrases.append(rng.choice([f"set the LoRA rank to {r}", f"use lora_r {r}"]))
if rng.random() < 0.5:
alpha = rng.choice([r, 2 * r])
fields["peft.lora_alpha"] = alpha
phrases.append(f"with alpha {alpha}")
if rng.random() < 0.6:
path = rng.choice(LOCAL_PATHS)
fields["data.train.datasets[0].dataset_name"] = "text_sft_jsonl"
fields["data.train.datasets[0].dataset_path"] = path
data_phrase = rng.choice([
f"on my conversation data at {path} (JSONL, oumi chat format)",
f"using the local JSONL file {path} as training data",
])
if rng.random() < 0.25:
val = re.sub(r"\.jsonl$", "_val.jsonl", path)
fields["data.validation.datasets[0].dataset_name"] = "text_sft_jsonl"
fields["data.validation.datasets[0].dataset_path"] = val
phrases.append(f"use {val} as the validation set")
else:
ds = rng.choice(HUB_DATASETS)
fields["data.train.datasets[0].dataset_name"] = ds
data_phrase = rng.choice([f"on the {ds} dataset", f"using {ds} from the Hub"])
if rng.random() < 0.2:
n = rng.choice([1000, 5000, 10000, 20000])
fields["data.train.datasets[0].sample_count"] = n
phrases.append(f"cap the training data at {n} examples")
if rng.random() < 0.2:
fields["data.train.pack"] = True
phrases.append("pack sequences")
if rng.random() < 0.65:
n = rng.choice([1, 2, 3, 4, 5])
fields["training.num_train_epochs"] = n
phrases.append(rng.choice([f"train for {n} epoch{'s' if n > 1 else ''}",
f"{n} epoch{'s' if n > 1 else ''}"]))
else:
n = rng.choice([100, 200, 500, 1000, 2000, 3000])
fields["training.max_steps"] = n
phrases.append(f"train for {n} steps")
lr = rng.choice(["1e-5", "2e-5", "5e-5", "1e-4", "2e-4", "3e-4", "5e-4"])
fields["training.learning_rate"] = float(lr)
phrases.append(rng.choice([f"learning rate {lr}", f"lr {lr}"]))
for prob, path, choices, phrase_fn in [
(0.6, "training.per_device_train_batch_size", [1, 2, 4, 8, 16],
lambda v: rng.choice([f"batch size {v} per device", f"per-device batch size of {v}"])),
(0.4, "training.gradient_accumulation_steps", [2, 4, 8, 16, 32],
lambda v: f"gradient accumulation {v}"),
(0.5, "training.mixed_precision_dtype", ["bf16", "fp16"],
lambda v: f"in {v} mixed precision"),
(0.4, "training.output_dir", [f"output/{s}" for s in PROJECT_SLUGS],
lambda v: rng.choice([f"write everything to {v}", f"output dir {v}"])),
(0.3, "training.seed", [0, 1, 7, 42, 123, 1234],
lambda v: f"seed {v}"),
(0.35, "training.lr_scheduler_type", ["cosine", "linear", "constant"],
lambda v: f"{v} LR schedule"),
(0.25, "training.warmup_steps", [10, 50, 100, 200],
lambda v: f"{v} warmup steps"),
(0.2, "training.weight_decay", [0.01, 0.1],
lambda v: f"weight decay {v}"),
(0.25, "training.save_steps", [100, 250, 500, 1000],
lambda v: f"checkpoint every {v} steps"),
(0.2, "training.logging_steps", [10, 25, 50, 100],
lambda v: f"log every {v} steps"),
(0.25, "training.enable_gradient_checkpointing", [True],
lambda v: "turn on gradient checkpointing"),
(0.3, "model.model_max_length", [1024, 2048, 4096, 8192],
lambda v: rng.choice([f"max sequence length {v}", f"context length {v}"])),
(0.2, "training.run_name", [f"{s}-{n}" for s in PROJECT_SLUGS for n in ("v1", "v2", "exp3")],
lambda v: f"call the run {v}"),
]:
if rng.random() < prob:
value = rng.choice(choices)
fields[path] = value
phrases.append(phrase_fn(value))
if rng.random() < 0.4:
n = rng.choice([50, 100, 200, 250, 500])
fields["training.eval_strategy"] = "steps"
fields["training.eval_steps"] = n
phrases.append(f"evaluate every {n} steps")
rng.shuffle(phrases)
opener = rng.choice(OPENERS).format(slug=rng.choice(PROJECT_SLUGS))
text = f"{opener}I want to {kind_phrase} {data_phrase}."
if phrases:
text += " " + ", ".join(phrases) + "."
return fields, text
def get_path(config: dict, path: str) -> object:
"""Navigate 'a.b[0].c'-style paths; returns None when absent."""
current: object = config
for part in re.findall(r"[^.\[\]]+|\[\d+\]", path):
if part.startswith("["):
if not isinstance(current, list):
return None
idx = int(part[1:-1])
if idx >= len(current):
return None
current = current[idx]
else:
if not isinstance(current, dict):
return None
current = current.get(part)
if current is None:
return None
return current
def values_match(expected: object, got: object) -> bool:
if isinstance(expected, bool) or isinstance(got, bool):
return isinstance(expected, bool) and isinstance(got, bool) and expected is got
if isinstance(expected, (int, float)) and isinstance(got, (int, float)):
return float(expected) == float(got)
if isinstance(expected, str) and isinstance(got, str):
return expected.strip().lower() == got.strip().lower()
return expected == got
# --- Task hooks ---------------------------------------------------------------
_rng = random.Random()
def sample() -> Sample:
seed = _rng.randrange(1, 10**9)
_, text = generate(seed)
return Sample(input_id=f"req-{seed}", text=text)
def lookup_truth(input_id: str) -> dict | None:
m = re.fullmatch(r"req-(\d+)", input_id or "")
if not m:
return None
fields, _ = generate(int(m.group(1)))
return {"fields": fields}
def parse_output(text: str) -> dict:
blocks = re.findall(r"```(?:ya?ml)?\s*\n(.*?)```", text, re.DOTALL)
yaml_text = blocks[-1].strip() if blocks else text.strip()
config, parse_error = None, None
try:
loaded = yaml.safe_load(yaml_text)
if isinstance(loaded, dict):
config = loaded
else:
parse_error = "config is not a YAML mapping"
except yaml.YAMLError as e:
parse_error = f"YAML parse error: {e}"
errors = [parse_error] if parse_error else validate(config)
return {"yaml_text": yaml_text, "config": config, "errors": errors}
def score(truth: dict, parsed: dict) -> dict[str, float]:
config = parsed["config"]
scores = {
"yaml_valid": float(config is not None),
"schema_valid": float(config is not None and not parsed["errors"]),
}
fields = truth["fields"]
matched = sum(
values_match(expected, get_path(config or {}, path))
for path, expected in fields.items()
)
scores["field_match"] = matched / len(fields)
return scores
def present(parsed: dict, truth: dict | None) -> dict:
out = {
"yaml": parsed["yaml_text"],
"yaml_valid": parsed["config"] is not None,
"schema_valid": parsed["config"] is not None and not parsed["errors"],
"errors": parsed["errors"],
}
if truth:
config = parsed["config"] or {}
out["fields"] = [
{
"path": path,
"expected": expected,
"got": get_path(config, path),
"match": values_match(expected, get_path(config, path)),
}
for path, expected in truth["fields"].items()
]
return out
TASK = Task(
order=2,
id="config-copilot",
title="Config Copilot",
tagline="Describe a training run, get a valid oumi config",
system_prompt=SYSTEM_PROMPT,
ui={
"output": "config",
"input_label": "Training run request",
"placeholder": "Deal a request, or describe the fine-tuning run you want...",
"deal_label": "Deal me a request",
},
sample=sample,
lookup_truth=lookup_truth,
parse_output=parse_output,
score=score,
present=present,
tools=TOOLS,
execute_tool=execute_tool,
# Routed to the "oumi" backend (an Oumi proxy deployment) when LLM_OUMI_* is
# configured; falls back to the default backend otherwise. The other tasks
# stay on the default backend. See config.backend() and .env.example.
backend="oumi",
)