ruhui 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- ruhui/__init__.py +59 -0
- ruhui/agent.py +385 -0
- ruhui/common.py +280 -0
- ruhui/email.py +104 -0
- ruhui/lang.py +292 -0
- ruhui/presets.py +187 -0
- ruhui/router.py +329 -0
- ruhui/shortlist.py +272 -0
- ruhui-0.1.0.dist-info/METADATA +175 -0
- ruhui-0.1.0.dist-info/RECORD +14 -0
- ruhui-0.1.0.dist-info/WHEEL +5 -0
- ruhui-0.1.0.dist-info/licenses/LICENSE +176 -0
- ruhui-0.1.0.dist-info/licenses/NOTICE +12 -0
- ruhui-0.1.0.dist-info/top_level.txt +1 -0
ruhui/__init__.py
ADDED
|
@@ -0,0 +1,59 @@
|
|
|
1
|
+
"""Ruhui: 非自回归 System 1 决策引擎(中文/多语言),带校准概率。
|
|
2
|
+
|
|
3
|
+
参照 Laya 架构 fork,命名取自房谋杜断的杜如晦(字克明),"晦"音近"hui",
|
|
4
|
+
寓"谋断"——System 1 快速决策。
|
|
5
|
+
"""
|
|
6
|
+
|
|
7
|
+
from .agent import Agent, RLAgent, load
|
|
8
|
+
from .common import (
|
|
9
|
+
QTYPES,
|
|
10
|
+
QTYPE_NAMES,
|
|
11
|
+
confidence_from_probs,
|
|
12
|
+
ece_score,
|
|
13
|
+
proper_reward,
|
|
14
|
+
render_options,
|
|
15
|
+
td_lambda_targets,
|
|
16
|
+
)
|
|
17
|
+
from .email import clean_email_body, email_state
|
|
18
|
+
from .lang import analyse as detect_language
|
|
19
|
+
from .lang import detect_script, is_english
|
|
20
|
+
from .presets import (
|
|
21
|
+
email_questions,
|
|
22
|
+
guard_questions,
|
|
23
|
+
moderation_questions,
|
|
24
|
+
router_questions,
|
|
25
|
+
triage_questions,
|
|
26
|
+
)
|
|
27
|
+
from .router import DEFAULT_MODELS, RouteDecision, Router
|
|
28
|
+
from .shortlist import embed_fn_from_agent, predict_shortlist, shortlist_choice
|
|
29
|
+
|
|
30
|
+
__version__ = "0.1.0"
|
|
31
|
+
__all__ = [
|
|
32
|
+
"Agent",
|
|
33
|
+
"RLAgent",
|
|
34
|
+
"load",
|
|
35
|
+
"Router",
|
|
36
|
+
"RouteDecision",
|
|
37
|
+
"DEFAULT_MODELS",
|
|
38
|
+
"shortlist_choice",
|
|
39
|
+
"predict_shortlist",
|
|
40
|
+
"embed_fn_from_agent",
|
|
41
|
+
"detect_language",
|
|
42
|
+
"detect_script",
|
|
43
|
+
"is_english",
|
|
44
|
+
"clean_email_body",
|
|
45
|
+
"email_questions",
|
|
46
|
+
"email_state",
|
|
47
|
+
"guard_questions",
|
|
48
|
+
"moderation_questions",
|
|
49
|
+
"router_questions",
|
|
50
|
+
"triage_questions",
|
|
51
|
+
"proper_reward",
|
|
52
|
+
"td_lambda_targets",
|
|
53
|
+
"ece_score",
|
|
54
|
+
"confidence_from_probs",
|
|
55
|
+
"render_options",
|
|
56
|
+
"QTYPES",
|
|
57
|
+
"QTYPE_NAMES",
|
|
58
|
+
"__version__",
|
|
59
|
+
]
|
ruhui/agent.py
ADDED
|
@@ -0,0 +1,385 @@
|
|
|
1
|
+
"""High-level inference runtime for ruhui System 1 decision models."""
|
|
2
|
+
import json
|
|
3
|
+
import os
|
|
4
|
+
import warnings
|
|
5
|
+
from typing import Any, Dict, Optional, Union
|
|
6
|
+
|
|
7
|
+
import numpy as np
|
|
8
|
+
import torch
|
|
9
|
+
|
|
10
|
+
from .common import (
|
|
11
|
+
QTYPES,
|
|
12
|
+
TEMP_MAX,
|
|
13
|
+
TEMP_MIN,
|
|
14
|
+
amp_dtype,
|
|
15
|
+
build_model,
|
|
16
|
+
build_sequence,
|
|
17
|
+
clamp_temperature,
|
|
18
|
+
collate_items,
|
|
19
|
+
confidence_from_probs,
|
|
20
|
+
render_options,
|
|
21
|
+
temp_bucket,
|
|
22
|
+
)
|
|
23
|
+
|
|
24
|
+
|
|
25
|
+
def _fix_tokenizer_config(path: str):
|
|
26
|
+
"""Ensure tokenizer_config.json can be loaded across all transformers versions."""
|
|
27
|
+
cfg_file = os.path.join(path, "tokenizer", "tokenizer_config.json")
|
|
28
|
+
if not os.path.exists(cfg_file):
|
|
29
|
+
return
|
|
30
|
+
try:
|
|
31
|
+
with open(cfg_file) as f:
|
|
32
|
+
tcfg = json.load(f)
|
|
33
|
+
changed = False
|
|
34
|
+
if tcfg.get("tokenizer_class") in (None, "TokenizersBackend"):
|
|
35
|
+
tcfg["tokenizer_class"] = "PreTrainedTokenizerFast"
|
|
36
|
+
tcfg.pop("backend", None)
|
|
37
|
+
tcfg.pop("is_local", None)
|
|
38
|
+
changed = True
|
|
39
|
+
# Checkpoints built on the mmBERT/Gemma tokenizer store extra_special_tokens as a list;
|
|
40
|
+
# transformers expects a mapping and raises "'list' object has no attribute 'keys'",
|
|
41
|
+
# which makes AutoTokenizer -- and so the whole model -- fail to load.
|
|
42
|
+
extra = tcfg.get("extra_special_tokens")
|
|
43
|
+
if isinstance(extra, list):
|
|
44
|
+
tcfg["extra_special_tokens"] = {"extra_%d" % i: t for i, t in enumerate(extra)}
|
|
45
|
+
changed = True
|
|
46
|
+
if changed:
|
|
47
|
+
with open(cfg_file, "w") as f:
|
|
48
|
+
json.dump(tcfg, f, indent=2)
|
|
49
|
+
except Exception:
|
|
50
|
+
pass
|
|
51
|
+
|
|
52
|
+
|
|
53
|
+
def _verify_compatibility(model: torch.nn.Module, cfg: Dict, weights: Dict[str, torch.Tensor], model_id: str):
|
|
54
|
+
"""Verify that the loaded checkpoint weights and config strictly match the expected architecture."""
|
|
55
|
+
# 1. Verify required configuration attributes
|
|
56
|
+
required_cfg = ["encoder", "head_layers"]
|
|
57
|
+
missing_cfg = [k for k in required_cfg if k not in cfg]
|
|
58
|
+
if missing_cfg:
|
|
59
|
+
raise ValueError(
|
|
60
|
+
f"Incompatible model config for {model_id!r}: missing configuration keys {missing_cfg}. "
|
|
61
|
+
f"Ensure this is a valid RL Agent decision model."
|
|
62
|
+
)
|
|
63
|
+
|
|
64
|
+
# 2. Check for required component prefixes
|
|
65
|
+
required_prefixes = ("encoder.", "type_emb.", "scorer.", "act_head.")
|
|
66
|
+
for prefix in required_prefixes:
|
|
67
|
+
if not any(k.startswith(prefix) for k in weights.keys()):
|
|
68
|
+
raise ValueError(
|
|
69
|
+
f"Incompatible model weights for {model_id!r}: checkpoint is missing '{prefix}' parameters. "
|
|
70
|
+
f"Expected an RL Agent decision model with encoder and decision heads."
|
|
71
|
+
)
|
|
72
|
+
|
|
73
|
+
# 3. Check for parameter shape mismatches
|
|
74
|
+
model_sd = model.state_dict()
|
|
75
|
+
shape_mismatches = []
|
|
76
|
+
missing_keys = []
|
|
77
|
+
|
|
78
|
+
for name, param in model.named_parameters():
|
|
79
|
+
if name not in weights:
|
|
80
|
+
missing_keys.append(name)
|
|
81
|
+
elif tuple(weights[name].shape) != tuple(param.shape):
|
|
82
|
+
shape_mismatches.append(f" - {name}: expected {tuple(param.shape)}, found {tuple(weights[name].shape)}")
|
|
83
|
+
|
|
84
|
+
if shape_mismatches:
|
|
85
|
+
err_details = "\n".join(shape_mismatches[:5])
|
|
86
|
+
if len(shape_mismatches) > 5:
|
|
87
|
+
err_details += f"\n ... and {len(shape_mismatches) - 5} more mismatched layers."
|
|
88
|
+
raise ValueError(
|
|
89
|
+
f"Model architecture mismatch for {model_id!r}:\n{err_details}\n"
|
|
90
|
+
f"The checkpoint weights do not match the configured model architecture."
|
|
91
|
+
)
|
|
92
|
+
|
|
93
|
+
if missing_keys:
|
|
94
|
+
raise ValueError(
|
|
95
|
+
f"Model weights incomplete for {model_id!r}: missing {len(missing_keys)} parameter tensors "
|
|
96
|
+
f"(e.g. {missing_keys[:3]})."
|
|
97
|
+
)
|
|
98
|
+
|
|
99
|
+
|
|
100
|
+
class Agent:
|
|
101
|
+
"""System 1 decision model runtime: fast, non-autoregressive, calibrated decisions."""
|
|
102
|
+
|
|
103
|
+
def __init__(
|
|
104
|
+
self,
|
|
105
|
+
model_id_or_path: str = "anyforge/ruhui",
|
|
106
|
+
device: Optional[str] = None,
|
|
107
|
+
token: Optional[str] = None,
|
|
108
|
+
subfolder: Optional[str] = None,
|
|
109
|
+
):
|
|
110
|
+
"""Load a Ruhui checkpoint.
|
|
111
|
+
|
|
112
|
+
`subfolder` selects one checkpoint from a repo that bundles several, e.g.
|
|
113
|
+
`Agent("anyforge/ruhui", subfolder="multilingual")`. Only that subfolder is
|
|
114
|
+
downloaded, so bundling does not cost every user the whole family.
|
|
115
|
+
"""
|
|
116
|
+
from safetensors.torch import load_file
|
|
117
|
+
from transformers import AutoTokenizer
|
|
118
|
+
|
|
119
|
+
model_dir = model_id_or_path
|
|
120
|
+
if not os.path.exists(model_dir):
|
|
121
|
+
if model_id_or_path.startswith(("/", "./", "../")) or os.path.isabs(model_id_or_path):
|
|
122
|
+
raise FileNotFoundError(
|
|
123
|
+
f"Local model path not found: {model_id_or_path!r}. "
|
|
124
|
+
f"Check that the directory exists and that training saved the model successfully."
|
|
125
|
+
)
|
|
126
|
+
from huggingface_hub import snapshot_download
|
|
127
|
+
|
|
128
|
+
# Restrict root checkpoints too: the default repo also contains sibling
|
|
129
|
+
# checkpoints, which an unfiltered snapshot would unnecessarily download.
|
|
130
|
+
prefix = f"{subfolder}/" if subfolder else ""
|
|
131
|
+
kw = {
|
|
132
|
+
"token": token or os.environ.get("HF_TOKEN"),
|
|
133
|
+
"allow_patterns": [prefix + name for name in (
|
|
134
|
+
"rl_agent_config.json", "model.safetensors", "tokenizer/*", "encoder/*",
|
|
135
|
+
)],
|
|
136
|
+
}
|
|
137
|
+
model_dir = snapshot_download(model_id_or_path, **kw)
|
|
138
|
+
|
|
139
|
+
if subfolder:
|
|
140
|
+
model_dir = os.path.join(model_dir, subfolder)
|
|
141
|
+
if not os.path.isdir(model_dir):
|
|
142
|
+
raise FileNotFoundError(
|
|
143
|
+
f"Subfolder {subfolder!r} not found in {model_id_or_path!r}."
|
|
144
|
+
)
|
|
145
|
+
|
|
146
|
+
_fix_tokenizer_config(model_dir)
|
|
147
|
+
|
|
148
|
+
cfg_path = os.path.join(model_dir, "rl_agent_config.json")
|
|
149
|
+
if not os.path.exists(cfg_path):
|
|
150
|
+
raise FileNotFoundError(
|
|
151
|
+
f"Incompatible model: {model_id_or_path!r} does not contain 'rl_agent_config.json'. "
|
|
152
|
+
f"That file ships with the weights of a Ruhui checkpoint, so load one of those "
|
|
153
|
+
f"(e.g. 'anyforge/ruhui') or a directory your own training run wrote."
|
|
154
|
+
)
|
|
155
|
+
|
|
156
|
+
with open(cfg_path) as f:
|
|
157
|
+
self.cfg = json.load(f)
|
|
158
|
+
|
|
159
|
+
weights_path = os.path.join(model_dir, "model.safetensors")
|
|
160
|
+
if not os.path.exists(weights_path):
|
|
161
|
+
raise FileNotFoundError(
|
|
162
|
+
f"Incompatible model: 'model.safetensors' not found in {model_id_or_path!r}."
|
|
163
|
+
)
|
|
164
|
+
|
|
165
|
+
# 1. Device resolution with automatic fallback
|
|
166
|
+
if device is not None:
|
|
167
|
+
target_device = torch.device(device)
|
|
168
|
+
if target_device.type == "cuda" and not torch.cuda.is_available():
|
|
169
|
+
print("Warning: CUDA requested but not available. Falling back to CPU.")
|
|
170
|
+
self.device = torch.device("cpu")
|
|
171
|
+
elif target_device.type == "mps" and not (hasattr(torch.backends, "mps") and torch.backends.mps.is_available()):
|
|
172
|
+
print("Warning: MPS requested but not available. Falling back to CPU.")
|
|
173
|
+
self.device = torch.device("cpu")
|
|
174
|
+
else:
|
|
175
|
+
self.device = target_device
|
|
176
|
+
else:
|
|
177
|
+
if torch.cuda.is_available():
|
|
178
|
+
self.device = torch.device("cuda")
|
|
179
|
+
elif hasattr(torch.backends, "mps") and torch.backends.mps.is_available():
|
|
180
|
+
self.device = torch.device("mps")
|
|
181
|
+
else:
|
|
182
|
+
self.device = torch.device("cpu")
|
|
183
|
+
|
|
184
|
+
tok_dir = os.path.join(model_dir, "tokenizer")
|
|
185
|
+
self.tok = AutoTokenizer.from_pretrained(tok_dir if os.path.exists(tok_dir) else self.cfg.get("encoder"))
|
|
186
|
+
|
|
187
|
+
enc_dir = os.path.join(model_dir, "encoder")
|
|
188
|
+
self.model = build_model(self.cfg, encoder_dir=enc_dir if os.path.exists(enc_dir) else None)
|
|
189
|
+
|
|
190
|
+
# Load weights and verify architectural compatibility
|
|
191
|
+
weights = load_file(weights_path)
|
|
192
|
+
_verify_compatibility(self.model, self.cfg, weights, model_id_or_path)
|
|
193
|
+
|
|
194
|
+
self.model.load_state_dict(weights, strict=True)
|
|
195
|
+
|
|
196
|
+
# ModernBERT's reference_compile defaults to "auto" and will torch.compile the encoder.
|
|
197
|
+
# That is a loss for the batch sizes Ruhui runs (a handful of questions per call) and can
|
|
198
|
+
# hang on some platforms, so keep the eager path.
|
|
199
|
+
try:
|
|
200
|
+
self.model.encoder.config.reference_compile = False
|
|
201
|
+
except Exception:
|
|
202
|
+
pass
|
|
203
|
+
|
|
204
|
+
# Keep what the checkpoint shipped for inspection, but only ever apply clamped values:
|
|
205
|
+
# some buckets are fitted to sharpen rather than soften (see clamp_temperature).
|
|
206
|
+
self.temperature_raw = self.cfg.get("temperature", [1.0, 1.0, 1.0])
|
|
207
|
+
self.temperature_by_options_raw = self.cfg.get("temperature_by_options", {})
|
|
208
|
+
self.temperature = [clamp_temperature(t) for t in self.temperature_raw]
|
|
209
|
+
self.temperature_by_options = {k: clamp_temperature(v)
|
|
210
|
+
for k, v in self.temperature_by_options_raw.items()}
|
|
211
|
+
rejected = ["%s=%.4g" % (k, float(v)) for k, v in self.temperature_by_options_raw.items()
|
|
212
|
+
if clamp_temperature(v) != float(v)]
|
|
213
|
+
rejected += ["temperature[%d]=%.4g" % (i, float(t)) for i, t in enumerate(self.temperature_raw)
|
|
214
|
+
if clamp_temperature(t) != float(t)]
|
|
215
|
+
if rejected:
|
|
216
|
+
warnings.warn(
|
|
217
|
+
"laya: this checkpoint ships temperatures outside [%g, %g] which would distort "
|
|
218
|
+
"confidence; clamping %s. Treat confidence from the affected buckets as uncalibrated."
|
|
219
|
+
% (TEMP_MIN, TEMP_MAX, ", ".join(rejected)),
|
|
220
|
+
RuntimeWarning, stacklevel=2)
|
|
221
|
+
self.dtype = amp_dtype(self.cfg.get("amp_dtype", "fp16"))
|
|
222
|
+
|
|
223
|
+
if self.device.type == "cuda" and torch.cuda.get_device_capability(self.device)[0] < 8:
|
|
224
|
+
self.dtype = torch.float16
|
|
225
|
+
elif self.device.type in ("cpu", "mps"):
|
|
226
|
+
self.dtype = torch.float32
|
|
227
|
+
|
|
228
|
+
# 2. Place on device with graceful fallback to CPU on memory error
|
|
229
|
+
fell_back_from = fell_back_why = None
|
|
230
|
+
try:
|
|
231
|
+
self.model.to(self.device).eval()
|
|
232
|
+
except (RuntimeError, torch.cuda.OutOfMemoryError) as e:
|
|
233
|
+
if self.device.type != "cpu":
|
|
234
|
+
# Record what actually went wrong: the reason matters more than the symptom,
|
|
235
|
+
# and it is the only place the underlying exception is ever surfaced.
|
|
236
|
+
fell_back_from, fell_back_why = self.device, e
|
|
237
|
+
self.device = torch.device("cpu")
|
|
238
|
+
self.dtype = torch.float32
|
|
239
|
+
self.model.to(self.device).eval()
|
|
240
|
+
else:
|
|
241
|
+
raise e
|
|
242
|
+
|
|
243
|
+
if fell_back_from is not None:
|
|
244
|
+
print(
|
|
245
|
+
"\n[ruhui] Warning: could not place the model on %s, so it is running on CPU.\n"
|
|
246
|
+
" Reason: %s\n"
|
|
247
|
+
" Inference will be roughly 10-15x slower (~200-500 ms rather than ~35 ms).\n"
|
|
248
|
+
" If this is a newer NVIDIA GPU (Blackwell / RTX 50-series), your PyTorch build\n"
|
|
249
|
+
" may not support its CUDA architecture:\n"
|
|
250
|
+
" pip install --pre torch --index-url https://download.pytorch.org/whl/nightly/cu128\n"
|
|
251
|
+
" See https://pytorch.org/get-started/locally/\n"
|
|
252
|
+
% (fell_back_from, fell_back_why), flush=True)
|
|
253
|
+
|
|
254
|
+
@staticmethod
|
|
255
|
+
def _to_internal(qdef: Dict) -> Dict:
|
|
256
|
+
t = qdef["type"]
|
|
257
|
+
crit = qdef.get("criteria")
|
|
258
|
+
if t == "choice" and isinstance(crit, list):
|
|
259
|
+
crit = {c: None for c in crit}
|
|
260
|
+
ins = qdef["instructions"]
|
|
261
|
+
if not isinstance(ins, str):
|
|
262
|
+
ins = json.dumps(ins)
|
|
263
|
+
return {"t": t, "ins": ins, "crit": crit}
|
|
264
|
+
|
|
265
|
+
@torch.no_grad()
|
|
266
|
+
def system_one(self, state: Union[str, dict, list], questions: Dict[str, Dict[str, Any]]) -> Dict[str, Any]:
|
|
267
|
+
"""Evaluate typed questions across state in a single, parallel forward pass.
|
|
268
|
+
|
|
269
|
+
Args:
|
|
270
|
+
state: Text string, JSON dict, or conversation turn list.
|
|
271
|
+
questions: Dictionary mapping question_id -> question definition.
|
|
272
|
+
- choice: {"type": "choice", "instructions": "...", "criteria": {"optA": "...", ...}}
|
|
273
|
+
- score: {"type": "score", "instructions": "...", "criteria": ["lvl0", "lvl1", ...]}
|
|
274
|
+
- noul: {"type": "noul", "instructions": "..."}
|
|
275
|
+
|
|
276
|
+
Returns:
|
|
277
|
+
Dictionary with answers, probabilities, calibrated confidence, and token usage.
|
|
278
|
+
"""
|
|
279
|
+
ids = list(questions.keys())
|
|
280
|
+
items = []
|
|
281
|
+
max_len = self.cfg.get("max_len", 512)
|
|
282
|
+
head_max_len = self.cfg.get("head_max_len", 192)
|
|
283
|
+
|
|
284
|
+
for qid in ids:
|
|
285
|
+
q = self._to_internal(questions[qid])
|
|
286
|
+
seq, markers = build_sequence(self.tok, state, q, max_len, head_max_len)
|
|
287
|
+
if len(markers) != len(render_options(q)):
|
|
288
|
+
raise ValueError("question %r options exceed head_max_len=%d" % (qid, head_max_len))
|
|
289
|
+
items.append({"ids": seq, "markers": markers, "qtype": QTYPES[q["t"]]})
|
|
290
|
+
|
|
291
|
+
b = collate_items([items], self.tok.pad_token_id)
|
|
292
|
+
use_amp = self.device.type == "cuda"
|
|
293
|
+
|
|
294
|
+
try:
|
|
295
|
+
with torch.autocast(device_type=self.device.type, dtype=self.dtype, enabled=use_amp):
|
|
296
|
+
logits, act = self.model(
|
|
297
|
+
b["input_ids"].to(self.device),
|
|
298
|
+
b["attention_mask"].to(self.device),
|
|
299
|
+
b["marker_pos"].to(self.device),
|
|
300
|
+
b["marker_mask"].to(self.device),
|
|
301
|
+
b["qtype"].to(self.device),
|
|
302
|
+
)
|
|
303
|
+
except (RuntimeError, torch.cuda.OutOfMemoryError) as e:
|
|
304
|
+
if self.device.type != "cpu" and ("memory" in str(e).lower() or "cuda" in str(e).lower()):
|
|
305
|
+
print("Warning: GPU memory exceeded during inference. Falling back to CPU...")
|
|
306
|
+
self.device = torch.device("cpu")
|
|
307
|
+
self.dtype = torch.float32
|
|
308
|
+
self.model.to(self.device)
|
|
309
|
+
logits, act = self.model(
|
|
310
|
+
b["input_ids"].to(self.device),
|
|
311
|
+
b["attention_mask"].to(self.device),
|
|
312
|
+
b["marker_pos"].to(self.device),
|
|
313
|
+
b["marker_mask"].to(self.device),
|
|
314
|
+
b["qtype"].to(self.device),
|
|
315
|
+
)
|
|
316
|
+
else:
|
|
317
|
+
raise e
|
|
318
|
+
|
|
319
|
+
logits = logits.float().cpu().numpy()
|
|
320
|
+
act = torch.softmax(act.float(), -1).cpu().numpy()
|
|
321
|
+
|
|
322
|
+
answers = {}
|
|
323
|
+
n_tokens = int(b["attention_mask"].sum())
|
|
324
|
+
|
|
325
|
+
for r, qid in enumerate(ids):
|
|
326
|
+
q = self._to_internal(questions[qid])
|
|
327
|
+
k = len(items[r]["markers"])
|
|
328
|
+
qt = QTYPES[q["t"]]
|
|
329
|
+
t_scale = self.temperature_by_options.get(temp_bucket(qt, k), self.temperature[qt])
|
|
330
|
+
z = logits[r, :k] / t_scale
|
|
331
|
+
p = np.exp(z - z.max())
|
|
332
|
+
p = p / p.sum()
|
|
333
|
+
|
|
334
|
+
conf_score = round(confidence_from_probs(p, k), 4)
|
|
335
|
+
ext = {"act_probability": round(float(act[r, 0]), 4)}
|
|
336
|
+
|
|
337
|
+
if q["t"] == "choice":
|
|
338
|
+
keys = list(q["crit"].keys())
|
|
339
|
+
answers[qid] = {
|
|
340
|
+
"type": "choice",
|
|
341
|
+
"choice": keys[int(p.argmax())],
|
|
342
|
+
"probabilities": {kk: round(float(v), 4) for kk, v in zip(keys, p)},
|
|
343
|
+
"confidence": conf_score,
|
|
344
|
+
"action": ext,
|
|
345
|
+
}
|
|
346
|
+
elif q["t"] == "score":
|
|
347
|
+
exp_score = float((np.arange(k) * p).sum())
|
|
348
|
+
answers[qid] = {
|
|
349
|
+
"type": "score",
|
|
350
|
+
"score": round(exp_score, 4),
|
|
351
|
+
"legend": {str(i): c for i, c in enumerate(q["crit"])},
|
|
352
|
+
"probabilities": {str(i): round(float(v), 4) for i, v in enumerate(p)},
|
|
353
|
+
"confidence": conf_score,
|
|
354
|
+
"action": ext,
|
|
355
|
+
}
|
|
356
|
+
else:
|
|
357
|
+
answers[qid] = {
|
|
358
|
+
"type": "noul",
|
|
359
|
+
"noul": round(float(p[1]), 4),
|
|
360
|
+
"confidence": round(max(float(p[1]), 1.0 - float(p[1])), 4),
|
|
361
|
+
"action": ext,
|
|
362
|
+
}
|
|
363
|
+
|
|
364
|
+
return {
|
|
365
|
+
"model": "ruhui-rl-agent",
|
|
366
|
+
"answers": answers,
|
|
367
|
+
"usage": {"input_tokens": n_tokens, "output_tokens": 0},
|
|
368
|
+
}
|
|
369
|
+
|
|
370
|
+
predict = system_one
|
|
371
|
+
|
|
372
|
+
|
|
373
|
+
RLAgent = Agent
|
|
374
|
+
|
|
375
|
+
|
|
376
|
+
def load(model_id_or_path: str = "anyforge/ruhui", device: Optional[str] = None,
|
|
377
|
+
token: Optional[str] = None, subfolder: Optional[str] = None) -> Agent:
|
|
378
|
+
"""Load a Ruhui agent.
|
|
379
|
+
|
|
380
|
+
`subfolder` picks one checkpoint out of a repo that bundles several:
|
|
381
|
+
|
|
382
|
+
ruhui.load("anyforge/ruhui") # English (repo root)
|
|
383
|
+
ruhui.load("anyforge/ruhui", subfolder="multilingual")
|
|
384
|
+
"""
|
|
385
|
+
return Agent(model_id_or_path, device=device, token=token, subfolder=subfolder)
|