#!/usr/bin/env python3
"""Tool-call SFT for BehzatOne v109 — runs ON Vast GPU only."""
from __future__ import annotations
import json
import os
import re
import shutil
from pathlib import Path
import torch
from datasets import concatenate_datasets, load_dataset
from peft import LoraConfig, PeftModel, get_peft_model
from transformers import AutoModelForCausalLM, AutoTokenizer
from trl import SFTConfig, SFTTrainer
BASE_REPO = os.environ.get("BASE_REPO", "behzatindustries/BehzatOne-8B-A1B")
BASE_SUB = os.environ.get("BASE_SUB", "v108")
OUT_REPO = os.environ.get("OUT_REPO", "behzatindustries/BehzatOne-8B-A1B")
OUT_SUB = os.environ.get("OUT_SUB", "v109")
HF_TOKEN = os.environ.get("HF_TOKEN", "")
WORK = Path(os.environ.get("WORK_DIR", "/workspace/behzatone-train"))
MAX_SAMPLES = int(os.environ.get("MAX_SAMPLES", "12000"))
MAX_STEPS = int(os.environ.get("MAX_STEPS", "1500"))
TOOL_CALL_RE = re.compile(r"\s*(\{.*?\})\s*", re.DOTALL)
FUNC_CALL_RE = re.compile(r"\s*(\{.*\})", re.DOTALL)
def patch_chat_template(tokenizer) -> None:
if getattr(tokenizer, "chat_template", None):
return
import urllib.request
url = f"https://huggingface.co/{BASE_REPO}/resolve/main/v106/chat_template.jinja"
req = urllib.request.Request(url, headers={"Authorization": f"Bearer {HF_TOKEN}"})
with urllib.request.urlopen(req, timeout=120) as r:
tokenizer.chat_template = r.read().decode("utf-8")
print("Loaded chat_template from v106")
def openai_tools_from_raw(raw) -> list[dict]:
if raw is None:
return []
if isinstance(raw, str):
raw = json.loads(raw)
out = []
for item in raw:
fn = item.get("function", item)
out.append(
{
"name": fn["name"],
"description": fn.get("description", ""),
"parameters": fn.get("parameters", {"type": "object", "properties": {}}),
}
)
return out
def parse_hermes_gpt(value: str) -> dict:
tool_calls = []
for match in TOOL_CALL_RE.finditer(value):
obj = json.loads(match.group(1))
args = obj.get("arguments", {})
if isinstance(args, str):
args = json.loads(args) if args.strip().startswith("{") else {"__raw__": args}
tool_calls.append(
{
"type": "function",
"function": {"name": obj["name"], "arguments": args},
}
)
content = TOOL_CALL_RE.sub("", value).strip()
if tool_calls:
return {"role": "assistant", "content": content or None, "tool_calls": tool_calls}
return {"role": "assistant", "content": value}
def hermes_row_to_messages(row: dict) -> tuple[list, list[dict]]:
tools = openai_tools_from_raw(row.get("tools"))
messages: list[dict] = []
for turn in row.get("conversations", []):
role = turn.get("from")
value = turn.get("value", "")
if role == "system":
continue
if role == "human":
messages.append({"role": "user", "content": value})
elif role == "gpt":
messages.append(parse_hermes_gpt(value))
elif role in {"tool", "function"}:
messages.append({"role": "tool", "content": value})
return messages, tools
def glaive_tools_from_system(system: str) -> list[dict]:
m = re.search(r"\{.*\}", system, re.DOTALL)
if not m:
return []
fn = json.loads(m.group(0))
return [
{
"name": fn["name"],
"description": fn.get("description", ""),
"parameters": fn.get("parameters", {"type": "object", "properties": {}}),
}
]
def glaive_row_to_messages(row: dict) -> tuple[list, list[dict]]:
tools = glaive_tools_from_system(row.get("system", ""))
messages: list[dict] = []
chat = row.get("chat", "")
for block in re.split(r"\n\n+", chat.strip()):
block = block.strip()
if not block:
continue
if block.startswith("USER:"):
messages.append({"role": "user", "content": block[5:].strip()})
elif block.startswith("ASSISTANT:"):
body = block[10:].strip().replace("<|endoftext|>", "").strip()
m = FUNC_CALL_RE.search(body)
if m:
obj = json.loads(m.group(1))
args = obj.get("arguments", "{}")
if isinstance(args, str):
args = json.loads(args) if args.strip().startswith("{") else {"__raw__": args}
messages.append(
{
"role": "assistant",
"content": None,
"tool_calls": [
{
"type": "function",
"function": {"name": obj["name"], "arguments": args},
}
],
}
)
else:
messages.append({"role": "assistant", "content": body})
elif block.startswith("FUNCTION RESPONSE:"):
messages.append({"role": "tool", "content": block[len("FUNCTION RESPONSE:") :].strip()})
return messages, tools
def row_to_text(tokenizer, row: dict, source: str) -> str:
if source == "hermes":
messages, tools = hermes_row_to_messages(row)
else:
messages, tools = glaive_row_to_messages(row)
if len(messages) < 2:
return ""
return tokenizer.apply_chat_template(
messages,
tools=tools or None,
tokenize=False,
add_generation_prompt=False,
)
def load_tool_datasets(tokenizer):
rows = []
hermes = load_dataset("NousResearch/hermes-function-calling-v1", split="train", token=HF_TOKEN)
hermes = hermes.shuffle(seed=42).select(range(min(7000, len(hermes))))
for row in hermes:
text = row_to_text(tokenizer, row, "hermes")
if "<|tool_call_start|>" in text:
rows.append({"text": text, "source": "hermes"})
glaive = load_dataset("glaiveai/glaive-function-calling-v2", split="train", token=HF_TOKEN)
glaive = glaive.shuffle(seed=42).select(range(min(5000, len(glaive))))
for row in glaive:
text = row_to_text(tokenizer, row, "glaive")
if "<|tool_call_start|>" in text:
rows.append({"text": text, "source": "glaive"})
from datasets import Dataset
ds = Dataset.from_list(rows)
if len(ds) > MAX_SAMPLES:
ds = ds.shuffle(seed=42).select(range(MAX_SAMPLES))
return ds
def main() -> None:
WORK.mkdir(parents=True, exist_ok=True)
status_path = WORK / "status.json"
def status(step: str, **extra):
payload = {"step": step, **extra}
status_path.write_text(json.dumps(payload, indent=2))
print("STATUS:", payload)
try:
from huggingface_hub import HfApi
HfApi(token=HF_TOKEN).upload_file(
path_or_fileobj=str(status_path),
path_in_repo="v109_train_status.json",
repo_id="rebehzat/behzatone-smoke-results",
repo_type="dataset",
commit_message=f"v109 status: {step}",
)
except Exception as exc:
print("status upload:", exc)
status("loading_tokenizer")
tokenizer = AutoTokenizer.from_pretrained(
BASE_REPO, subfolder=BASE_SUB, trust_remote_code=True, token=HF_TOKEN
)
patch_chat_template(tokenizer)
if tokenizer.pad_token is None:
tokenizer.pad_token = tokenizer.eos_token
status("loading_model")
model = AutoModelForCausalLM.from_pretrained(
BASE_REPO,
subfolder=BASE_SUB,
torch_dtype=torch.bfloat16,
trust_remote_code=True,
token=HF_TOKEN,
device_map="auto",
)
lora = LoraConfig(
r=32,
lora_alpha=64,
lora_dropout=0.05,
target_modules=["q_proj", "k_proj", "v_proj", "o_proj", "gate_proj", "up_proj", "down_proj"],
bias="none",
task_type="CAUSAL_LM",
)
model = get_peft_model(model, lora)
model.print_trainable_parameters()
status("loading_dataset")
ds = load_tool_datasets(tokenizer)
print(f"Training samples with native tool tokens: {len(ds)}")
if len(ds) < 100:
raise SystemExit("Too few tool-call training rows after conversion")
out_dir = WORK / "lora-out"
args = SFTConfig(
output_dir=str(out_dir),
per_device_train_batch_size=1,
gradient_accumulation_steps=8,
learning_rate=1.5e-5,
max_steps=MAX_STEPS,
warmup_ratio=0.05,
lr_scheduler_type="cosine",
logging_steps=10,
save_steps=500,
bf16=True,
max_length=4096,
dataset_text_field="text",
packing=False,
report_to="none",
)
status("training", samples=len(ds), max_steps=MAX_STEPS)
trainer = SFTTrainer(model=model, args=args, train_dataset=ds, processing_class=tokenizer)
trainer.train()
trainer.save_model(str(out_dir / "final"))
status("merging")
base = AutoModelForCausalLM.from_pretrained(
BASE_REPO,
subfolder=BASE_SUB,
torch_dtype=torch.bfloat16,
trust_remote_code=True,
token=HF_TOKEN,
device_map="cpu",
)
merged = PeftModel.from_pretrained(base, str(out_dir / "final"))
merged = merged.merge_and_unload()
merged_dir = WORK / "merged-v109"
merged_dir.mkdir(exist_ok=True)
merged.save_pretrained(merged_dir, safe_serialization=True)
tokenizer.save_pretrained(merged_dir)
import urllib.request
tpl_url = f"https://huggingface.co/{BASE_REPO}/resolve/main/v106/chat_template.jinja"
req = urllib.request.Request(tpl_url, headers={"Authorization": f"Bearer {HF_TOKEN}"})
with urllib.request.urlopen(req, timeout=120) as r:
(merged_dir / "chat_template.jinja").write_bytes(r.read())
status("uploading")
from huggingface_hub import HfApi
api = HfApi(token=HF_TOKEN)
api.upload_folder(
folder_path=str(merged_dir),
path_in_repo=OUT_SUB,
repo_id=OUT_REPO,
repo_type="model",
commit_message="v109: tool-call SFT on v108 (Hermes + Glaive, native LFM2.5 tool tokens)",
)
status("done", out=f"{OUT_REPO}/{OUT_SUB}")
print("TRAINING COMPLETE")
if __name__ == "__main__":
main()