jebadiah-9b-v2 / scripts /decide_standalone.py
jbrashear's picture
Jebadiah 9B v2: model card, merge verification, standalone script
35debf7 verified
Raw History Blame
2.77 kB
"""Run Jebadiah standalone: render a request exactly as AINode's /v1/systemone does, read the label-token
logits at the answer position (fp32), apply the model's per-type temperatures, print typed answers.
python decide_standalone.py --model <this repository, a local dir or frontier-infra/jebadiah-9b-v2> --request req.json
v2 ships merged bf16 weights, so there is no base or adapter to combine: the model directory carries the
weights, the tokenizer, the chat template and temperatures.json. req.json is a /v1/systemone body:
{"state": ..., "questions": {id: {"type", "instructions", "criteria"}}}. The three modules beside this file
are the trainer's own copies of the renderer (ainode_prompt_verbatim.py is AINode's decide rendering copied
verbatim at commit e5c08938) and the logit read, so the bytes match training.
"""
import argparse, json, os, sys
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
import torch
from jebadiah_model import load_tokenizer, load_base, read_temperatures, Scorer
from jebadiah_prompt import answer_from_probs
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--model", required=True, help="the merged model: a local directory or a Hub repo id")
ap.add_argument("--revision", default=None, help="Hub revision, when --model is a repo id")
ap.add_argument("--request", required=True, help="JSON file with state and questions")
ap.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu")
ap.add_argument("--no-temperatures", action="store_true", help="raw probabilities, as the served route returns today")
ap.add_argument("--attn", default="sdpa")
a = ap.parse_args()
req = json.load(open(a.request))
if a.device.startswith("mps"):
# transformers 5.17's threaded weight loader intermittently segfaults moving tensors to MPS
os.environ.setdefault("HF_DEACTIVATE_ASYNC_LOAD", "1")
dtype = torch.bfloat16 if a.device.startswith(("cuda", "mps")) else torch.float32
model_dir = a.model
if not os.path.isdir(model_dir):
from huggingface_hub import snapshot_download
model_dir = snapshot_download(a.model, revision=a.revision)
tok = load_tokenizer(model_dir)
model = load_base(model_dir, attn_implementation=a.attn, dtype=dtype, device=a.device)
temps = {} if a.no_temperatures else read_temperatures(model_dir)
scorer = Scorer(model, tok, temperatures=temps, device=a.device)
res = scorer.score(req["state"], req["questions"])
out = {"temperatures_applied": temps, "answers": {}}
for qid, (keys, probs) in res.items():
out["answers"][qid] = answer_from_probs(req["questions"][qid], keys, probs)
print(json.dumps(out, indent=1))
if __name__ == "__main__":
main()