Spaces:
Sleeping
Sleeping
File size: 8,195 Bytes
796da7c | 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 | from __future__ import annotations
import argparse
import copy
import json
from pathlib import Path
from typing import Any, Dict, List
from unsloth import FastLanguageModel
from permanence.agent_interface.parser import parse_agent_output
from permanence.env import PermanenceEnv
from training.config import TrainingConfig, load_simple_yaml
DEFAULT_CONFIG_PATH = "training/config.yaml"
DEFAULT_STATE_PATH = Path("dashboard") / "current_state.json"
DEFAULT_GHOST_RECORDING_PATH = Path("ghost_recording.json")
MAX_NEW_TOKENS = 220
MAX_SEQ_LENGTH = 2048
CASCADE_PLAN = [
{
"action_id": "review_contract_internally",
"completion": (
"<thinking>Start with internal review to preserve downstream options and prevent premature lockouts.</thinking>\n"
'<action id="review_contract_internally" contract_id="cascade_contract_001"/>\n'
'<reversibility level="R1" confidence="0.97"/>'
),
},
{
"action_id": "align_with_legal",
"completion": (
"<thinking>Legal alignment is required before external communication to keep amendments valid.</thinking>\n"
'<action id="align_with_legal" dispute_summary="Internal review complete; legal terms aligned."/>\n'
'<reversibility level="R2" confidence="0.93"/>'
),
},
{
"action_id": "communicate_resolution_externally",
"completion": (
"<thinking>Now communicate externally with aligned terms to resolve without triggering cascade locks.</thinking>\n"
'<action id="communicate_resolution_externally" client_id="client_a" resolution_terms="Aligned remediation and amended timeline" final_amount="1500"/>\n'
'<reversibility level="R3" confidence="0.91"/>'
),
},
]
def _resolve_model_dir(config_path: str, model_path: str | None) -> Path:
if model_path:
return Path(model_path)
config_data = load_simple_yaml(config_path)
config = TrainingConfig.from_mapping(config_data)
return Path(config.output_dir) / "final_model"
def _load_trained_model(model_dir: Path):
if not model_dir.exists():
raise FileNotFoundError(f"Trained model not found at {model_dir}")
model, tokenizer = FastLanguageModel.from_pretrained(
model_name=str(model_dir),
max_seq_length=MAX_SEQ_LENGTH,
dtype=None,
load_in_4bit=True,
)
if tokenizer.pad_token is None:
tokenizer.pad_token = tokenizer.eos_token
if hasattr(FastLanguageModel, "for_inference"):
try:
model = FastLanguageModel.for_inference(model)
except Exception:
pass
return model, tokenizer
def _generate_candidate_completion(model, tokenizer, prompt: str, max_new_tokens: int) -> str:
inputs = tokenizer(prompt, return_tensors="pt")
device = getattr(model, "device", None)
if device is not None:
inputs = {key: value.to(device) for key, value in inputs.items()}
output_ids = model.generate(
**inputs,
max_new_tokens=max_new_tokens,
do_sample=True,
temperature=0.7,
top_p=0.9,
eos_token_id=tokenizer.eos_token_id,
pad_token_id=tokenizer.pad_token_id,
)
generated = output_ids[:, inputs["input_ids"].shape[1] :]
return tokenizer.decode(generated[0], skip_special_tokens=True)
def _build_prompt(observation_text: str, expected_action_id: str) -> str:
return (
"You are solving PERMANENCE Task 5 (Cascade).\n"
"Return strictly: <thinking>...</thinking> then one <action id=\"...\" .../> and one <reversibility level=\"R1-R5\" confidence=\"0-1\"/>.\n"
f"Prioritize action id: {expected_action_id}.\n\n"
f"Observation:\n{observation_text}\n"
)
def _build_dashboard_payload(env: PermanenceEnv, episode_data: Dict[str, Any], metrics: Dict[str, Any]) -> Dict[str, Any]:
ws = env._current_world_state
if ws is None:
raise RuntimeError("World state is missing")
recent_actions = []
for record in ws.action_history[-5:]:
recent_actions.append(
{
"action": record.action_id,
"r_level": record.actual_r_level,
"step": record.step,
"predicted_r_level": record.predicted_r_level,
"predicted_confidence": record.predicted_confidence,
}
)
return {
"recent_actions": recent_actions,
"locked_actions": dict(ws.locked_actions),
"critical_options": dict(ws.critical_options),
"catastrophe_rate": metrics.get("recent_catastrophe_rate", []),
"episode": metrics.get("total_episodes", 0),
"episode_data": episode_data,
"raw_thinking": str(episode_data.get("raw_thinking", "")),
}
def run_ghost_export(model, tokenizer, state_path: Path, recording_path: Path) -> Dict[str, Any]:
env = PermanenceEnv(config={"force_task": "task_cascade"})
observation, info = env.reset(seed=12345)
metrics: Dict[str, Any] = {"total_episodes": 1, "recent_catastrophe_rate": []}
timeline: List[Dict[str, Any]] = []
state_path.parent.mkdir(parents=True, exist_ok=True)
for index, planned_step in enumerate(CASCADE_PLAN, start=1):
prompt = _build_prompt(observation.get("text", ""), planned_step["action_id"])
candidate = _generate_candidate_completion(model, tokenizer, prompt, max_new_tokens=MAX_NEW_TOKENS)
parsed_candidate = parse_agent_output(candidate)
completion = candidate
if parsed_candidate.action_id != planned_step["action_id"]:
completion = planned_step["completion"]
parsed_final = parse_agent_output(completion)
observation, reward, terminated, truncated, step_info = env.step(completion)
catastrophe = 1.0 if step_info.get("reward_breakdown", {}).get("catastrophe_count", 0) > 0 else 0.0
rates = list(metrics.get("recent_catastrophe_rate", []))
rates.append(catastrophe)
metrics["recent_catastrophe_rate"] = rates[-50:]
episode_data = {
"prompt": prompt,
"completion": completion,
"observation": observation,
"reward": float(reward),
"terminated": bool(terminated),
"truncated": bool(truncated),
"info": step_info,
"raw_thinking": parsed_final.raw_thinking or "",
"step_index": index,
"task_id": info.get("task_id", "task_cascade"),
}
payload = _build_dashboard_payload(env, episode_data, metrics)
state_path.write_text(json.dumps(payload, indent=2), encoding="utf-8")
timeline.append(copy.deepcopy(payload))
if terminated or truncated:
break
recording_path.write_text(json.dumps(timeline, indent=2), encoding="utf-8")
final_reason = ""
if timeline:
final_reason = str(timeline[-1].get("episode_data", {}).get("info", {}).get("termination_reason", ""))
if final_reason != "success":
raise RuntimeError(
f"Task 5 ghost export did not complete successfully (termination_reason={final_reason or 'none'})"
)
return {
"steps_recorded": len(timeline),
"recording_path": str(recording_path),
"state_path": str(state_path),
"termination_reason": final_reason,
}
def main() -> None:
parser = argparse.ArgumentParser(description="Export offline ghost demo recording for dashboard playback")
parser.add_argument("--config", default=DEFAULT_CONFIG_PATH)
parser.add_argument("--model-path", default=None)
parser.add_argument("--state-path", default=str(DEFAULT_STATE_PATH))
parser.add_argument("--output", default=str(DEFAULT_GHOST_RECORDING_PATH))
args = parser.parse_args()
model_dir = _resolve_model_dir(args.config, args.model_path)
model, tokenizer = _load_trained_model(model_dir)
summary = run_ghost_export(
model=model,
tokenizer=tokenizer,
state_path=Path(args.state_path),
recording_path=Path(args.output),
)
print(json.dumps(summary, indent=2))
if __name__ == "__main__":
main() |