Dot-Loom-Conductor-14B / inference_example.py
usedot's picture
Publish Dot Loom Conductor 14B adapter and evidence
43bff50 verified
Raw
History Blame Contribute Delete
3.87 kB
from __future__ import annotations
import json
import os
import torch
from peft import PeftModel
from transformers import Mistral3ForConditionalGeneration, MistralCommonBackend
BASE_ID = os.environ.get("DOT_BASE_MODEL", "mistralai/Ministral-3-14B-Base-2512")
BASE_REVISION = os.environ.get(
"DOT_BASE_REVISION", "5b0ceedbb42dff466ae60b258ba296f32da51384"
)
ADAPTER_ID = os.environ.get("DOT_CONDUCTOR_ADAPTER", "usedot/Dot-Loom-Conductor-14B")
task = {
"family": "security_review",
"category": "security",
"summary": "Audit a user-controlled fetcher for SSRF and metadata access.",
"risk": 0.95,
"complexity": 0.88,
"consequence": 0.96,
"ambiguity": 0.62,
"evidence_need": 0.92,
"reversibility": 0.18,
"input_tokens": 2100,
}
constraints = {
"max_calls": 3,
"max_credits": 16.5,
"max_latency_ms": 95000,
"minimum_quality": 0.90,
}
workers = [
{
"id": "worker_a",
"provider_group": "local",
"quality": 0.72,
"pass_rate": 0.56,
"credits_per_call": 0.15,
"p95_latency_ms": 6500,
"strengths": ["coding", "drafting", "speed"],
},
{
"id": "worker_b",
"provider_group": "provider_y",
"quality": 0.89,
"pass_rate": 0.85,
"credits_per_call": 3.5,
"p95_latency_ms": 35000,
"strengths": ["reasoning", "review", "writing"],
},
{
"id": "worker_c",
"provider_group": "provider_z",
"quality": 0.97,
"pass_rate": 0.96,
"credits_per_call": 12.5,
"p95_latency_ms": 42000,
"strengths": ["implementation", "reasoning", "synthesis"],
},
]
schema = {
"policy": "lean|balanced|strict",
"writer": "worker id",
"reviewer": "worker id|null",
"finalizer": "worker id|null",
"max_calls": "integer",
"max_credits": "number",
"max_latency_ms": "integer",
"estimated_credits": "number",
"estimated_latency_ms": "integer",
"estimated_quality": "number",
"estimated_pass_rate": "number",
"independent_verification": "boolean",
"access": {"writer": [], "reviewer": [], "finalizer": []},
"reason_codes": ["bounded machine-readable codes"],
}
payload = {"task": task, "constraints": constraints, "workers": workers}
prompt = (
"You are the Dot Loom conductor. Select the highest-utility execution plan that obeys all hard "
"call, credit, and latency limits. Use cheap workers when sufficient. Require an independent "
"provider when consequences and verification needs justify it. Return exactly one compact JSON "
"object and no prose.\nINPUT="
+ json.dumps(payload, sort_keys=True, separators=(",", ":"))
+ "\nOUTPUT_SCHEMA="
+ json.dumps(schema, sort_keys=True, separators=(",", ":"))
+ "\nOUTPUT="
)
tokenizer = MistralCommonBackend.from_pretrained(BASE_ID, revision=BASE_REVISION)
base = Mistral3ForConditionalGeneration.from_pretrained(
BASE_ID,
revision=BASE_REVISION,
dtype=torch.bfloat16,
device_map="auto",
attn_implementation="sdpa",
)
model = PeftModel.from_pretrained(base, ADAPTER_ID).eval()
token_ids = tokenizer.encode(prompt, add_special_tokens=False)
if hasattr(token_ids, "tolist"):
token_ids = token_ids.tolist()
if token_ids and isinstance(token_ids[0], list):
token_ids = token_ids[0]
input_ids = torch.tensor([token_ids], dtype=torch.long, device=model.device)
eos_id = int(tokenizer.eos_token_id)
with torch.inference_mode():
output = model.generate(
input_ids=input_ids,
attention_mask=torch.ones_like(input_ids),
max_length=input_ids.shape[1] + 260,
do_sample=False,
eos_token_id=eos_id,
pad_token_id=eos_id,
use_cache=True,
)
generated = output[0, input_ids.shape[1] :].tolist()
print(tokenizer.decode(generated, skip_special_tokens=True))