Instructions to use vishwr/claim_drafter with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- PEFT
How to use vishwr/claim_drafter with PEFT:
from peft import PeftModel from transformers import AutoModelForCausalLM base_model = AutoModelForCausalLM.from_pretrained("Qwen/Qwen3.5-9B") model = PeftModel.from_pretrained(base_model, "vishwr/claim_drafter") - Notebooks
- Google Colab
- Kaggle
File size: 12,458 Bytes
164e036 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 | #!/usr/bin/env python3
"""Verify the whole setup end-to-end before spending the training budget.
python3 scripts/preflight.py
Runs nine checks, cheapest first, stopping at the first hard failure. Checks 1-5
are free and offline. Checks 6-9 hit the Tinker API and cost a few cents in
total: one training step on 2 examples, and one short generation.
What it proves, in order:
1 .env present and TINKER_API_KEY loaded
2 tinker + tinker_cookbook importable at the expected versions
3 datasets present and well-formed
4 reward function still calibrated (catches a broken refactor)
5 no example exceeds max_length (silent right-truncation would eat claims)
6 API key authenticates and the target model is actually served
7 renderer masks the prompt and trains only on the assistant tokens
8 forward_backward + optim_step complete and return a finite loss
9 the model generates, and the generation is scorable
Exit code 0 means you are clear to run training/train_sft.py.
"""
import json
import os
import sys
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
from claim_drafter.config import DEFAULTS, load_env
OK, BAD, WARN = " [PASS]", " [FAIL]", " [WARN]"
_state = {"failed": 0, "warned": 0}
def check(name):
print("\n%s" % name)
def ok(msg):
print("%s %s" % (OK, msg))
def warn(msg):
_state["warned"] += 1
print("%s %s" % (WARN, msg))
def fail(msg, fatal=True):
_state["failed"] += 1
print("%s %s" % (BAD, msg))
if fatal:
print("\nPreflight stopped. Fix the above and re-run.")
sys.exit(1)
AUTH_URL = ("https://tinker.thinkingmachines.dev/services/tinker-prod"
"/api/v1/auth/token")
def _probe_auth_directly(api_key, base_url=None):
"""Ask the auth endpoint whether the key is good, using stdlib urllib.
Returns "valid", "rejected", or a short description of the transport failure.
Deliberately does NOT use the tinker SDK: the point is to take a different
network path so we can tell a bad key apart from a blocked connection.
"""
import urllib.error
import urllib.request
url = (base_url.rstrip("/") + "/api/v1/auth/token") if base_url else AUTH_URL
request = urllib.request.Request(
url, data=b"{}", method="POST",
headers={"X-API-Key": api_key, "Content-Type": "application/json",
"User-Agent": "claim-drafter-preflight/1.0"},
)
try:
with urllib.request.urlopen(request, timeout=45) as response:
return "valid" if response.status == 200 else "http %d" % response.status
except urllib.error.HTTPError as e:
return "rejected" if e.code in (401, 403) else "http %d" % e.code
except Exception as e:
return "%s: %s" % (type(e).__name__, str(e)[:80])
def main():
model = os.environ.get("CLAIM_DRAFTER_MODEL", DEFAULTS["model"])
print("=" * 68)
print("Claim drafter preflight model=%s" % model)
print("=" * 68)
# ---------------------------------------------------------------- 1
check("1. Environment")
if not os.path.exists(".env"):
fail(".env not found. Run: cp .env.example .env and add your key.")
load_env()
key = os.environ["TINKER_API_KEY"]
if len(key) < 8:
fail("TINKER_API_KEY looks too short to be real.")
ok("TINKER_API_KEY loaded (%s…%s, %d chars)" % (key[:4], key[-2:], len(key)))
# ---------------------------------------------------------------- 2
check("2. Packages")
try:
import tinker
ok("tinker %s" % getattr(tinker, "__version__", "(version attr missing)"))
except ImportError:
fail("tinker not installed. pip install -r requirements.txt")
try:
import tinker_cookbook # noqa: F401
ok("tinker_cookbook importable")
except ImportError:
fail("tinker_cookbook not installed. pip install -r requirements.txt")
# ---------------------------------------------------------------- 3
check("3. Datasets")
train_path = "data/sft/train.jsonl"
if not os.path.exists(train_path):
fail("%s missing. Rebuild with: make data (see docs/pipeline.md)" % train_path)
with open(train_path) as f:
rows = [json.loads(l) for l in f if l.strip()]
ok("%s: %d examples" % (train_path, len(rows)))
roles = [m["role"] for m in rows[0]["messages"]]
if roles != ["system", "user", "assistant"]:
fail("unexpected role sequence %s; expected system/user/assistant" % roles)
ok("role sequence system/user/assistant")
dpo_path = "data/dpo/comparisons.jsonl"
if os.path.exists(dpo_path):
n = sum(1 for _ in open(dpo_path))
ok("%s: %d preference pairs" % (dpo_path, n))
else:
warn("%s missing — stages 2/3 unavailable until you run make dpo" % dpo_path)
# ---------------------------------------------------------------- 4
check("4. Reward function calibration")
from claim_drafter.rewards import claim_reward
scores = [claim_reward(r["messages"][2]["content"]) for r in rows[:300]]
mean = sum(scores) / len(scores)
if mean < 0.95:
fail("mean reward on real granted claims is %.3f; expected >0.95. "
"The checker is broken, not the claims." % mean)
ok("mean reward on 300 real granted claim sets: %.3f" % mean)
broken = claim_reward("1. A thing.\n\n3. The thing of claim 9.")
if broken >= mean:
fail("reward does not punish malformed claims (%.3f vs %.3f)" % (broken, mean))
ok("malformed claim set scores lower (%.3f)" % broken)
# ---------------------------------------------------------------- 5
check("5. Sequence lengths")
from tinker_cookbook.tokenizer_utils import get_tokenizer
tokenizer = get_tokenizer(model)
longest = 0
for r in rows:
n = sum(len(tokenizer.encode(m["content"])) for m in r["messages"])
longest = max(longest, n)
if longest + 32 > DEFAULTS["max_length"]:
fail("longest example is %d tokens vs max_length %d. The cookbook "
"truncates silently FROM THE RIGHT, which would cut off the end of "
"the claim set." % (longest, DEFAULTS["max_length"]))
ok("longest example %d tokens, under max_length %d — nothing truncates"
% (longest, DEFAULTS["max_length"]))
# ---------------------------------------------------------------- 6
check("6. API authentication and model availability")
import tinker
try:
service_client = tinker.ServiceClient()
caps = service_client.get_server_capabilities()
served = [m.model_name for m in caps.supported_models]
ok("authenticated; server lists %d models" % len(served))
except Exception as e:
# Distinguish "bad key" from "cannot reach Tinker". The SDK reports both
# as APIConnectionError, which is genuinely confusing. Re-test the auth
# endpoint with stdlib urllib: it takes a different network path than the
# SDK's HTTP client, so a 200 here means the key is fine and the problem
# is transport (corporate proxy, VPN, sandboxed shell).
verdict = _probe_auth_directly(key)
if verdict == "valid":
fail("Your API KEY IS VALID, but the tinker SDK could not connect.\n"
" This is a transport problem, not a credentials problem.\n"
" Do NOT set TINKER_BASE_URL -- the default endpoint is correct.\n"
" The usual cause is pyqwest 0.7.0, which builds an EMPTY CA root\n"
" store and fails every TLS handshake. requirements.txt pins\n"
" pyqwest<0.7.0 and claim_drafter/transport.py repairs an already-\n"
" installed 0.7.0 at runtime. If you still see this, check that\n"
" load_env() ran, and try TINKER_TRANSPORT=httpx.\n"
" Underlying error: %s" % e)
elif verdict == "rejected":
fail("TINKER_API_KEY was REJECTED by the server (HTTP 401).\n"
" Check the value in .env against your Tinker dashboard.")
else:
fail("Could not reach Tinker at all (%s).\n"
" Underlying error: %s" % (verdict, e))
if model not in served:
fail("%s is not in the served model list. Available Qwen models: %s"
% (model, [m for m in served if "Qwen" in m][:8]))
ok("%s is served" % model)
# ---------------------------------------------------------------- 7
check("7. Renderer and loss masking")
from tinker_cookbook import renderers
from tinker_cookbook.renderers import TrainOnWhat
from tinker_cookbook.supervised.data import conversation_to_datum
renderer = renderers.get_renderer(DEFAULTS["renderer"], tokenizer)
ok("renderer %s" % DEFAULTS["renderer"])
datum = conversation_to_datum(
rows[0]["messages"], renderer, DEFAULTS["max_length"],
train_on_what=TrainOnWhat.LAST_ASSISTANT_MESSAGE,
)
weights = datum.loss_fn_inputs["weights"]
weights = getattr(weights, "data", weights)
weights = list(weights)
nonzero = sum(1 for w in weights if w > 0)
if nonzero == 0:
fail("every loss weight is zero — nothing would be trained on")
if nonzero == len(weights):
fail("every loss weight is non-zero — the prompt is not masked")
ok("%d/%d tokens carry loss (prompt masked, assistant trained)"
% (nonzero, len(weights)))
# ---------------------------------------------------------------- 8
check("8. One training step (costs a few cents)")
training_client = service_client.create_lora_training_client(
base_model=model, rank=DEFAULTS["lora_rank"])
ok("LoRA training client created (rank %d)" % DEFAULTS["lora_rank"])
batch = [
conversation_to_datum(r["messages"], renderer, DEFAULTS["max_length"],
train_on_what=TrainOnWhat.LAST_ASSISTANT_MESSAGE)
for r in rows[:2]
]
n_tokens = sum(d.model_input.length for d in batch)
# Submit both, then await both, so the server pipelines them.
fwd_bwd_future = training_client.forward_backward(batch, loss_fn="cross_entropy")
optim_future = training_client.optim_step(
tinker.AdamParams(learning_rate=1e-4, beta1=0.9, beta2=0.95, eps=1e-8))
fwd_bwd_result = fwd_bwd_future.result()
optim_future.result()
ok("forward_backward + optim_step completed on %d tokens" % n_tokens)
try:
logprobs = fwd_bwd_result.loss_fn_outputs[0]["logprobs"]
vals = list(getattr(logprobs, "data", logprobs))
if not all(v == v for v in vals[:50]): # NaN check
fail("logprobs contain NaN")
ok("logprobs finite")
except (KeyError, IndexError, TypeError):
warn("could not read logprobs from the result; step still succeeded")
# ---------------------------------------------------------------- 9
check("9. Sampling")
# No `name=`: it is deprecated and has no effect -- these checkpoints are
# always ephemeral. Use save_weights_for_sampler(name=...) for a durable one.
sampling_client = training_client.save_weights_and_get_sampling_client()
prompt = renderer.build_generation_prompt(rows[0]["messages"][:2])
result = sampling_client.sample(
prompt=prompt,
sampling_params=tinker.types.SamplingParams(
max_tokens=256, temperature=0.0, stop=renderer.get_stop_sequences()),
num_samples=1,
).result()
message, _ = renderer.parse_response(result.sequences[0].tokens)
text = renderers.get_text_content(message)
if not text.strip():
fail("model returned an empty generation")
ok("generated %d chars" % len(text))
print("\n ---- first 300 chars ----")
for line in text[:300].splitlines():
print(" | %s" % line)
print(" -------------------------")
warn("untrained LoRA output will look poor — this only proves the path works")
print("\n" + "=" * 68)
if _state["failed"]:
print("PREFLIGHT FAILED (%d)" % _state["failed"])
sys.exit(1)
print("PREFLIGHT PASSED%s" % (" with %d warning(s)" % _state["warned"]
if _state["warned"] else ""))
print("Next: python3 training/train_sft.py --epochs 2")
print("=" * 68)
if __name__ == "__main__":
main()
|