载入中...
搜索中...
未找到
AgentBindings.cpp
浏览该文件的文档.
1#include "agent/AgentModule.h"
2#include "agent/Codec.h"
4
5#include <algorithm>
6#include <exception>
7
8namespace eve::agent {
9namespace {
10struct StackScope {
12 SQInteger top;
13 explicit StackScope(HSQUIRRELVM v) : vm(v), top(sq_gettop(v)) {}
14 ~StackScope() { sq_settop(vm, top); }
15};
16struct ActiveScope {
17 bool& active;
18 explicit ActiveScope(bool& value) : active(value) { active = true; }
19 ~ActiveScope() { active = false; }
20};
21
22// Script references live only for this synchronous adapter call. Every callback
23// restores the VM stack, including missing methods and thrown script errors.
24class ScriptEnvironment final : public IEnvironment {
25public:
26 explicit ScriptEnvironment(const ssq::Table& environment) : environment_(environment) {}
27 Result<Observation> reset(std::uint64_t seed) override { return invoke("reset", seed, 0, 0); }
28 Result<Observation> step(std::uint32_t action, double dt) override { return invoke("step", 0, action, dt); }
29
30private:
31 Result<Observation> invoke(const char* method, std::uint64_t seed, std::uint32_t action, double dt) {
32 auto vm = environment_.getHandle();
33 StackScope stack(vm);
34 sq_pushobject(vm, environment_.getRaw());
35 sq_pushstring(vm, method, -1);
36 if (SQ_FAILED(sq_get(vm, -2)))
38 Diagnostic::error(DiagnosticCode::NotFound, "Missing environment method", {}, {}, "agent.script"));
39 sq_pushobject(vm, environment_.getRaw());
40 const bool resetting = std::string_view(method) == "reset";
41 if (resetting) {
42 const auto text = std::to_string(seed);
43 sq_pushstring(vm, text.c_str(), SQInteger(text.size()));
44 } else {
45 sq_pushinteger(vm, SQInteger(action));
46 sq_pushfloat(vm, SQFloat(dt));
47 }
48 if (SQ_FAILED(sq_call(vm, resetting ? 2 : 3, SQTrue, SQFalse))) {
49 sq_getlasterror(vm);
50 const SQChar* text = nullptr;
51 const auto converted = sq_getstring(vm, -1, &text);
53 DiagnosticCode::CallbackFailure, SQ_SUCCEEDED(converted) && text ? text : "Environment callback failed",
54 {}, {}, "agent.script"));
55 }
56 auto value = script::valueFromSquirrel(vm, -1);
57 if (!value) return Result<Observation>::failure(value.status());
58 return decodeObservation(value.value());
59 }
60 ssq::Table environment_;
61};
62
63ssq::Table project(HSQUIRRELVM vm, Result<Value> result) {
64 // The common projector has a bounded payload. Check first so oversized
65 // reports return a failure instead of an ok result with a null payload.
66 if (result) {
67 StackScope stack(vm);
68 auto checked = script::pushValue(vm, result.value());
69 if (!checked)
70 return script::projectResult(vm, Result<Value>::failure(checked.status()), [](Value v) { return v; });
71 }
72 return script::projectResult(vm, std::move(result), [](Value v) { return v; });
73}
74Result<Value> inferScript(const ssq::Object& policy, const ssq::Object& observation, Backend backend, bool actionOnly) {
75 auto pv = script::valueFromSquirrel(policy);
76 if (!pv) return Result<Value>::failure(pv.status());
77 auto ov = script::valueFromSquirrel(observation);
78 if (!ov) return Result<Value>::failure(ov.status());
79 auto p = decodePolicy(pv.value());
80 if (!p) return Result<Value>::failure(p.status());
81 auto o = decodeObservation(ov.value());
82 if (!o) return Result<Value>::failure(o.status());
83 auto result = infer(p.value(), o.value(), backend);
84 if (!result) return Result<Value>::failure(result.status());
85 if (actionOnly) {
86 const auto& probabilities = result.value();
87 const auto action = *std::max_element(o.value().legalActions.begin(), o.value().legalActions.end(),
88 [&](auto a, auto b) { return probabilities[a] < probabilities[b]; });
89 return Result<Value>::success(Value(std::int64_t(action)));
90 }
91 Value::Array values;
92 for (auto n : result.value()) values.emplace_back(n);
93 return Result<Value>::success(Value(std::move(values)));
94}
95} // namespace
96
97void Agent::expose(ssq::Class& cls) {
98 const auto vm = cls.getHandle();
99 cls.addFunc("getName", &Agent::getName);
100 cls.addFunc("run", [vm](Agent* self, ssq::Table config, ssq::Table environment) {
101 if (self->active_)
102 return project(vm, Result<Value>::failure(Diagnostic::error(
103 DiagnosticCode::Conflict, "Agent run/replay reentry", {}, {}, "agent.script")));
104 ActiveScope active(self->active_);
105 auto value = script::valueFromSquirrel(config);
106 if (!value) return project(vm, Result<Value>::failure(value.status()));
107 auto parsed = decodeConfig(value.value());
108 if (!parsed) return project(vm, Result<Value>::failure(parsed.status()));
109 ScriptEnvironment adapter(environment);
110 auto result = run(parsed.value(), adapter);
111 if (!result) return project(vm, Result<Value>::failure(result.status()));
112 return project(vm, Result<Value>::success(encodeReport(result.value())));
113 });
114 cls.addFunc("infer", [vm](Agent*, ssq::Table policy, ssq::Table observation) {
115 return project(vm, inferScript(policy, observation, Backend::Cpu, false));
116 });
117 cls.addFunc("inferTensor", [vm](Agent*, ssq::Table policy, ssq::Table observation) {
118 return project(vm, inferScript(policy, observation, Backend::Tensor, false));
119 });
120 cls.addFunc("inferGpu", [vm](Agent*, ssq::Table policy, ssq::Table observation) {
121 return project(vm, inferScript(policy, observation, Backend::Gpu, false));
122 });
123 cls.addFunc("act", [vm](Agent*, ssq::Table policy, ssq::Table observation) {
124 return project(vm, inferScript(policy, observation, Backend::Cpu, true));
125 });
126 cls.addFunc("actTensor", [vm](Agent*, ssq::Table policy, ssq::Table observation) {
127 return project(vm, inferScript(policy, observation, Backend::Tensor, true));
128 });
129 cls.addFunc("actGpu", [vm](Agent*, ssq::Table policy, ssq::Table observation) {
130 return project(vm, inferScript(policy, observation, Backend::Gpu, true));
131 });
132 cls.addFunc("replay", [vm](Agent* self, ssq::Table trace, ssq::Table environment, float tolerance) {
133 if (self->active_)
134 return script::projectResult(
135 vm, Result<void>::failure(Diagnostic::error(DiagnosticCode::Conflict, "Agent run/replay reentry", {},
136 {}, "agent.script")));
137 ActiveScope active(self->active_);
138 auto value = script::valueFromSquirrel(trace);
139 if (!value) return script::projectResult(vm, Result<void>::failure(value.status()));
140 auto parsed = decodeTrace(value.value());
141 if (!parsed) return script::projectResult(vm, Result<void>::failure(parsed.status()));
142 ScriptEnvironment adapter(environment);
143 return script::projectResult(vm, replay(parsed.value(), adapter, tolerance));
144 });
145 cls.addFunc("getBackends", [vm](Agent*) {
146 Value::Array names{Value("cpu-mlp")};
147 auto tensor = backendName(Backend::Tensor);
148 if (tensor) names.emplace_back(tensor.value());
149 auto gpu = backendName(Backend::Gpu);
150 if (gpu) names.emplace_back(gpu.value());
151 return project(vm, Result<Value>::success(Value(std::move(names))));
152 });
153}
154} // namespace eve::agent
double value
HSQUIRRELVM vm
SQInteger top
bool & active
Trace trace
Definition Agent.cpp:49
struct SQVM * HSQUIRRELVM
glm::vec4 p[6]
HSQUIRRELVM vm
Definition ECS.cpp:20
HSQOBJECT cls
Definition ECS.cpp:21
std::map< std::string, Var > values
std::uint16_t method
glm::vec3 n
Definition Grass.cpp:63
float v
std::string text
MeleePoint3 b
Definition MeleeHit.cpp:41
MeleePoint3 a
Definition MeleeHit.cpp:40
std::weak_ptr< Run > run
Definition OnnxGpgpu.cpp:25
const RuntimeTensor & tensor
Definition OnnxLstm.cpp:26
std::string action
Definition PlayHost.cpp:117
std::uint32_t seed
Definition PointSet.cpp:807
The single Squirrel projection for common Result, Status and Value.
std::string_view adapter
float step
Definition TreeMesh.cpp:314
static Diagnostic error(DiagnosticCode code, std::string message, std::string path={}, DiagnosticDetails details={}, std::string source={})
Construct an error diagnostic with the standard error severity.
Definition Diagnostic.h:125
static Result failure(Status status)
Construct a failed result from a structured status.
Definition Result.h:175
Result< std::vector< double > > infer(const Policy &policy, const Observation &observation, Backend backend)
Compute masked action probabilities with validated version-1 owning weights.
Definition Agent.cpp:112
Result< std::string > backendName(Backend backend)
Return an owning backend label, or Unsupported if Tensor is unavailable; owner thread only.
Definition Agent.cpp:93
Result< Trace > decodeTrace(const Value &value)
Decode version-1 owning replay data; run-time consistency is checked before replay reset.
Definition Codec.cpp:207
Result< Policy > decodePolicy(const Value &value)
Decode and validate version-1 owning policy; unknown keys/versions rejected atomically.
Definition Codec.cpp:189
Result< void > replay(const Trace &trace, IEnvironment &environment, double tolerance)
Reset and replay actual actions, checking all observations and failure evidence.
Definition Agent.cpp:258
Result< Observation > decodeObservation(const Value &value)
Decode owning observation data, rejecting malformed values; dimensions are checked by run/infer.
Definition Codec.cpp:186
Value encodeReport(const Report &r)
Encode runner-produced report, including policy and replayable evidence, as owning data.
Definition Codec.cpp:247
Result< Config > decodeConfig(const Value &value)
Decode strict script/JSON configuration; unknown keys rejected, no mutations or callbacks.
Definition Codec.cpp:117
std::variant< std::monostate, std::int64_t, double, std::string, bool > Value
Definition Database.h:26
ssq::Table project(HSQUIRRELVM vm, const eve::Result< void > &result)
Project a completed editing result that has no payload.