载入中...
搜索中...
未找到
Agent.cpp
浏览该文件的文档.
1#include "agent/Agent.h"
2#include "agent/Learning.h"
3#include "common/Capability.h"
4
5#include <algorithm>
6#include <cmath>
7#include <limits>
8#include <set>
9
10namespace eve::agent {
11namespace {
12
13bool validObservation(const Observation& o, std::size_t features, std::size_t actions) {
14 if (o.features.size() != features || o.legalActions.size() > actions || !std::isfinite(o.reward) ||
15 std::abs(o.reward) > 1e9 || o.coverage.size() > 64 || o.finding.size() > 4096)
16 return false;
17 if (o.outcome != Outcome::Running && o.outcome != Outcome::Success && o.outcome != Outcome::Failure) return false;
18 if (o.outcome == Outcome::Running && o.legalActions.empty()) return false;
19 if (o.outcome == Outcome::Failure && o.finding.empty()) return false;
20 std::set<std::uint32_t> seen;
21 for (auto a : o.legalActions)
22 if (a >= actions || !seen.insert(a).second) return false;
23 for (float f : o.features)
24 if (!std::isfinite(f) || std::abs(f) > 1e6f) return false;
25 for (const auto& key : o.coverage)
26 if (key.empty() || key.size() > 256) return false;
27 return true;
28}
29
30bool validConfig(const Config& c) {
31 return c.featureCount > 0 && c.featureCount <= 1024 && c.actionCount > 0 && c.actionCount <= 1024 &&
32 c.hiddenWidth > 0 && c.hiddenWidth <= 64 && c.population > 0 && c.population <= 256 && c.generations > 0 &&
33 c.generations <= 1024 && c.horizon > 0 && c.horizon <= 1024 &&
34 std::uint64_t(c.population) * c.generations * c.horizon <= 1000000 &&
35 std::uint64_t(c.population) * c.horizon * c.featureCount <= 1048576 &&
36 (std::uint64_t(c.population) + c.maxFindings + 1) * c.horizon <= 16384 &&
37 std::uint64_t(c.generations) * c.trainingEpochs * c.elites * c.horizon <= 1000000 && c.elites > 0 &&
38 c.elites <= c.population && c.trainingEpochs <= 64 && c.maxFindings <= 256 && std::isfinite(c.dt) &&
39 c.dt > 0 && c.dt <= 60 && std::isfinite(c.mutationProbability) && c.mutationProbability >= 0 &&
40 c.mutationProbability <= 1 && std::isfinite(c.randomProbability) && c.randomProbability >= 0 &&
41 c.randomProbability <= 1 && std::isfinite(c.coverageWeight) && c.coverageWeight >= 0 &&
42 c.coverageWeight <= 1e6 && std::isfinite(c.failureWeight) && c.failureWeight >= 0 &&
43 c.failureWeight <= 1e6 && std::isfinite(c.learningRate) && c.learningRate > 0 && c.learningRate <= 1 &&
44 (c.strategy == Strategy::Random || c.strategy == Strategy::EvolutionLearning) &&
45 (c.backend == Backend::Cpu || c.backend == Backend::Tensor || c.backend == Backend::Gpu);
46}
47
48struct Candidate {
49 Trace trace;
50 double score = 0;
51};
52
53Result<std::uint32_t> sample(const Policy& policy, const Observation& o, detail::Random& random, bool uniform,
54 Backend backend) {
55 if (uniform) return Result<std::uint32_t>::success(o.legalActions[random.index(o.legalActions.size())]);
56 auto evaluated = infer(policy, o, backend);
57 if (!evaluated) return Result<std::uint32_t>::failure(evaluated.status());
58 const auto& probabilities = evaluated.value();
59 double u = random.unit();
60 for (auto a : o.legalActions) {
61 u -= probabilities[a];
62 if (u <= 0) return Result<std::uint32_t>::success(a);
63 }
64 return Result<std::uint32_t>::success(o.legalActions.back());
65}
66
67bool matches(const Observation& a, const Observation& b, double tolerance) {
68 if (a.features.size() != b.features.size() || a.legalActions != b.legalActions || a.coverage != b.coverage ||
69 a.outcome != b.outcome || a.finding != b.finding || !std::isfinite(b.reward) ||
70 std::abs(a.reward - b.reward) > tolerance)
71 return false;
72 for (std::size_t i = 0; i < a.features.size(); ++i)
73 if (!std::isfinite(b.features[i]) || std::abs(double(a.features[i]) - b.features[i]) > tolerance) return false;
74 return true;
75}
76
77} // namespace
78
80 if (policy.schemaId != "evengine.agent.policy" || policy.schemaVersion != 1 || policy.featureCount == 0 ||
81 policy.featureCount > 1024 || policy.actionCount == 0 || policy.actionCount > 1024 || policy.hiddenWidth == 0 ||
82 policy.hiddenWidth > 64 ||
83 policy.weights.size() != detail::weightCount(policy.featureCount, policy.hiddenWidth, policy.actionCount))
85 DiagnosticCode::InvalidArgument, "Unsupported policy schema, version or dimensions", {}, {}, "agent"));
86 for (double w : policy.weights)
87 if (!std::isfinite(w) || std::abs(w) > 1e6)
89 Diagnostic::error(DiagnosticCode::InvalidArgument, "Invalid policy weight", {}, {}, "agent"));
90 return Result<void>::success();
91}
92
94 if (backend == Backend::Cpu) return Result<std::string>::success("cpu-mlp");
95 if (backend == Backend::Gpu) {
96 if (auto* provider = cap::query<IGpuPolicyBackend>()) {
97 auto ready = provider->available();
98 if (!ready) return Result<std::string>::failure(ready.status());
99 return Result<std::string>::success(provider->name());
100 }
102 DiagnosticCode::Unsupported, "GPU requires an active AgentTensor module", {}, {}, "agent"));
103 }
104 if (backend != Backend::Tensor)
106 Diagnostic::error(DiagnosticCode::InvalidArgument, "Unknown backend", {}, {}, "agent"));
107 if (auto* provider = cap::query<IPolicyBackend>()) return Result<std::string>::success(provider->name());
109 DiagnosticCode::Unsupported, "Tensor backend requires an active AgentTensor module", {}, {}, "agent"));
110}
111
112Result<std::vector<double>> infer(const Policy& policy, const Observation& observation, Backend backend) {
113 auto validated = validatePolicy(policy);
114 if (!validated) return Result<std::vector<double>>::failure(validated.status());
115 if (!validObservation(observation, policy.featureCount, policy.actionCount) || observation.legalActions.empty())
117 DiagnosticCode::InvalidArgument, "Invalid observation or empty legal action mask", {}, {}, "agent"));
118 if (backend == Backend::Cpu) return Result<std::vector<double>>::success(detail::forward(policy, observation));
119 auto available = backendName(backend);
120 if (!available) return Result<std::vector<double>>::failure(available.status());
121 IPolicyBackend* provider = backend == Backend::Gpu ? cap::query<IGpuPolicyBackend>() : cap::query<IPolicyBackend>();
122 EV_ASSERT(provider, "backend registration must remain stable during inference");
123 auto result = provider->evaluate(policy, observation);
124 if (!result) return result;
125 if (result.value().size() != policy.actionCount)
126 return Result<std::vector<double>>::failure(
127 Diagnostic::error(DiagnosticCode::InvalidArgument, "Invalid backend shape", {}, {}, "agent"));
128 double sum = 0;
129 for (std::size_t i = 0; i < result.value().size(); ++i) {
130 const auto value = result.value()[i];
131 if (!std::isfinite(value) || value < 0 ||
132 (std::find(observation.legalActions.begin(), observation.legalActions.end(), i) ==
133 observation.legalActions.end() &&
134 value != 0))
135 return Result<std::vector<double>>::failure(
136 Diagnostic::error(DiagnosticCode::InvalidArgument, "Invalid backend probability", {}, {}, "agent"));
137 sum += value;
138 }
139 if (std::abs(sum - 1) > 1e-5)
141 DiagnosticCode::InvalidArgument, "Backend probabilities must sum to one", {}, {}, "agent"));
142 return result;
143}
144
145Result<Report> run(const Config& c, IEnvironment& environment) {
146 if (!validConfig(c))
148 "Invalid dimensions, budgets, strategy, probabilities or time",
149 {}, {}, "agent"));
150 auto selectedBackend = backendName(c.backend);
151 if (!selectedBackend) return Result<Report>::failure(selectedBackend.status());
152 Report report;
153 report.policy = detail::makePolicy(c);
154 report.backend = c.strategy == Strategy::Random ? "uniform-random" : selectedBackend.value();
155 report.trainingBackend = c.strategy == Strategy::Random ? "none"
156 : c.backend == Backend::Gpu ? "tensor-gpu-sgd"
157 : "cpu-sgd";
158 report.bestScore = -std::numeric_limits<double>::infinity();
159 detail::Random random(c.searchSeed);
160 std::vector<Candidate> parents;
161 std::set<std::string> coverage, findingKeys;
162 for (std::uint32_t generation = 0; generation < c.generations; ++generation) {
163 std::vector<Candidate> candidates;
164 for (std::uint32_t member = 0; member < c.population; ++member) {
165 Candidate candidate;
166 candidate.trace.environmentSeed = c.environmentSeed;
167 candidate.trace.dt = c.dt;
168 auto reset = environment.reset(c.environmentSeed);
169 if (!reset) return Result<Report>::failure(reset.status());
170 if (!validObservation(reset.value(), c.featureCount, c.actionCount) || reset.value().reward != 0)
172 "Environment returned invalid initial observation", {},
173 {}, "agent"));
174 candidate.trace.initial = std::move(reset).takeValue();
175 Observation observation = candidate.trace.initial;
176 std::set<std::string> episodeCoverage(observation.coverage.begin(), observation.coverage.end());
177 const auto parent = parents.empty() ? 0 : random.index(parents.size());
178 for (std::uint32_t tick = 0; tick < c.horizon && observation.outcome == Outcome::Running; ++tick) {
179 const bool uniform = c.strategy == Strategy::Random || random.unit() < c.randomProbability;
180 auto sampled = sample(report.policy, observation, random, uniform, c.backend);
181 if (!sampled) return Result<Report>::failure(sampled.status());
182 auto action = sampled.value();
183 if (!uniform && !parents.empty() && tick < parents[parent].trace.steps.size() &&
184 random.unit() >= c.mutationProbability) {
185 const auto inherited = parents[parent].trace.steps[tick].action;
186 if (std::find(observation.legalActions.begin(), observation.legalActions.end(), inherited) !=
187 observation.legalActions.end())
188 action = inherited;
189 }
190 auto stepped = environment.step(action, c.dt);
191 if (!stepped) return Result<Report>::failure(stepped.status());
192 if (!validObservation(stepped.value(), c.featureCount, c.actionCount))
194 "Environment returned invalid step observation",
195 {}, {}, "agent"));
196 observation = std::move(stepped).takeValue();
197 candidate.score += observation.reward;
198 episodeCoverage.insert(observation.coverage.begin(), observation.coverage.end());
199 candidate.trace.steps.push_back({action, observation});
200 ++report.steps;
201 }
202 ++report.episodes;
203 coverage.insert(episodeCoverage.begin(), episodeCoverage.end());
204 if (coverage.size() > 65536)
206 DiagnosticCode::Cancelled, "Coverage storage budget exceeded (65536 points)", {}, {}, "agent"));
207 candidate.score += c.coverageWeight * double(episodeCoverage.size());
208 if (observation.outcome == Outcome::Failure) {
209 ++report.failures;
210 candidate.score += c.failureWeight;
211 if (report.findings.size() < c.maxFindings && findingKeys.insert(observation.finding).second)
212 report.findings.push_back(candidate.trace);
213 }
214 if (candidate.score > report.bestScore) {
215 report.bestScore = candidate.score;
216 report.best = candidate.trace;
217 }
218 candidates.push_back(std::move(candidate));
219 }
220 std::stable_sort(candidates.begin(), candidates.end(),
221 [](const auto& a, const auto& b) { return a.score > b.score; });
222 candidates.resize(c.elites);
223 if (c.strategy == Strategy::EvolutionLearning) {
224 for (std::uint32_t epoch = 0; epoch < c.trainingEpochs; ++epoch)
225 for (const auto& elite : candidates) {
226 const Observation* previous = &elite.trace.initial;
227 for (const auto& step : elite.trace.steps) {
228 if (c.backend == Backend::Gpu) {
229 auto* provider = cap::query<IGpuPolicyBackend>();
230 if (!provider)
233 "GPU provider was removed during an environment callback", {}, {}, "agent"));
234 auto trained = provider->train(report.policy, *previous, step.action, c.learningRate);
235 if (!trained) return Result<Report>::failure(trained.status());
236 auto valid = validatePolicy(trained.value());
237 if (!valid) return Result<Report>::failure(valid.status());
238 if (trained.value().featureCount != c.featureCount ||
239 trained.value().hiddenWidth != c.hiddenWidth ||
240 trained.value().actionCount != c.actionCount)
242 "GPU training changed policy shape",
243 {}, {}, "agent"));
244 report.policy = std::move(trained).takeValue();
245 } else
246 detail::train(report.policy, *previous, step.action, c.learningRate);
247 previous = &step.observation;
248 ++report.trainingSamples;
249 }
250 }
251 parents = std::move(candidates);
252 }
253 }
254 report.coverage.assign(coverage.begin(), coverage.end());
255 return Result<Report>::success(std::move(report));
256}
257
258Result<void> replay(const Trace& trace, IEnvironment& environment, double tolerance) {
259 if (trace.schemaId != "evengine.agent.trace" || trace.schemaVersion != 1 || !std::isfinite(trace.dt) ||
260 trace.dt <= 0 || trace.dt > 60 || !std::isfinite(tolerance) || tolerance < 0 || trace.steps.size() > 1024 ||
261 trace.initial.features.empty() || trace.initial.features.size() > 1024 || trace.initial.reward != 0)
263 "Unsupported trace schema/version, time, length or tolerance",
264 {}, {}, "agent"));
266 if (!validObservation(*previous, trace.initial.features.size(), 1024))
268 Diagnostic::error(DiagnosticCode::InvalidArgument, "Invalid initial trace observation", {}, {}, "agent"));
269 for (const auto& step : trace.steps) {
270 if (previous->outcome != Outcome::Running ||
271 std::find(previous->legalActions.begin(), previous->legalActions.end(), step.action) ==
272 previous->legalActions.end() ||
273 !validObservation(step.observation, trace.initial.features.size(), 1024))
275 "Invalid trace action or observation", {}, {}, "agent"));
276 previous = &step.observation;
277 }
278 auto reset = environment.reset(trace.environmentSeed);
279 if (!reset) return Result<void>::failure(reset.status());
280 if (!matches(trace.initial, reset.value(), tolerance)) return Result<void>::failure(
281 Diagnostic::error(DiagnosticCode::Conflict, "Replay observation diverged", {}, {}, "agent"));
282 for (const auto& step : trace.steps) {
283 auto result = environment.step(step.action, trace.dt);
284 if (!result) return Result<void>::failure(result.status());
285 if (!matches(step.observation, result.value(), tolerance)) return Result<void>::failure(
286 Diagnostic::error(DiagnosticCode::Conflict, "Replay observation diverged", {}, {}, "agent"));
287 }
288 return Result<void>::success();
289}
290
291} // namespace eve::agent
double value
double score
Definition Agent.cpp:50
Trace trace
Definition Agent.cpp:49
float w
Definition AnimClip.cpp:738
#define EV_ASSERT(cond,...)
Assert an internal engine invariant (state that must always hold).
Definition Assert.h:37
std::uint32_t key
float u
Definition Grass.cpp:233
std::int32_t c
std::int32_t parent
bool valid
MeleePoint3 b
Definition MeleeHit.cpp:41
MeleePoint3 a
Definition MeleeHit.cpp:40
graphics::Canvas * previous
std::weak_ptr< Run > run
Definition OnnxGpgpu.cpp:25
uint64_t epoch
Definition OnnxGpgpu.cpp:28
float f
std::string action
Definition PlayHost.cpp:117
std::vector< ActionSpec > actions
Definition PlayHost.cpp:126
std::uint32_t generation
SimulationTick tick
Battle::Random random
float step
Definition TreeMesh.cpp:314
float size
Definition TreeMesh.cpp:156
float(ui::Theme::* member)[4]
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
Move-only operation result carrying either a value or Status.
Definition Result.h:155
static Result success(T value)
Construct a successful result owning value.
Definition Result.h:164
static Result failure(Status status)
Construct a failed result from a structured status.
Definition Result.h:175
Adapter for a resettable game, simulation, UI or other decision-making environment....
Definition Agent.h:36
virtual Result< Observation > reset(std::uint64_t seed)=0
Reset atomically to the injected seed; return initial state (reward zero).
virtual Result< Observation > step(std::uint32_t action, double dt)=0
Execute one legal action, advance exactly dt seconds, then inspect invariants.
Optional inference service. Providers own their registration and revoke before destruction....
Definition Agent.h:114
virtual Result< std::vector< double > > evaluate(const Policy &policy, const Observation &observation)=0
Evaluate inputs validated by agent::infer; return masked probabilities or structured failure.
Random public API.
Definition Learning.h:14
Policy makePolicy(const Config &c)
Make policy.
Definition Learning.h:77
void train(Policy &p, const Observation &o, std::uint32_t action, double rate)
Train.
Definition Learning.h:90
std::vector< double > forward(const Policy &p, const Observation &o)
Forward.
Definition Learning.h:65
std::size_t weightCount(std::size_t inputs, std::size_t hidden, std::size_t actions)
Weight count.
Definition Learning.h:34
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< void > validatePolicy(const Policy &policy)
Validate the complete owning policy before inference/import, with no mutation or callbacks.
Definition Agent.cpp:79
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
Backend
Explicit backend; Tensor is eager CPU, Gpu accelerates inference and SGD via agent_tensor.
Definition Agent.h:50
double sample(const Heightmap &map, double u, double v)
Sample.
Bounded search configuration; seed streams for environment, search and learning are separate.
Definition Agent.h:53
Owning state projection; action IDs index a fixed domain action catalogue.
Definition Agent.h:16
std::vector< std::uint32_t > legalActions
Definition Agent.h:18
std::string finding
Definition Agent.h:22
std::vector< std::string > coverage
Definition Agent.h:19
std::vector< float > features
Definition Agent.h:17
Owning version-1 portable network weights; import validates the entire value before use.
Definition Agent.h:99
std::uint32_t hiddenWidth
Definition Agent.h:104
std::string schemaId
Definition Agent.h:100
std::uint32_t schemaVersion
Definition Agent.h:101
std::uint32_t actionCount
Definition Agent.h:103
std::vector< double > weights
Definition Agent.h:105
std::uint32_t featureCount
Definition Agent.h:102
Owning search result; reward, coverage and failures remain separate evidence.
Definition Agent.h:148
std::vector< std::string > coverage
Definition Agent.h:152
double bestScore
Definition Agent.h:157
std::vector< Trace > findings
Definition Agent.h:151
std::uint64_t failures
Definition Agent.h:155
std::string trainingBackend
Definition Agent.h:159
std::string backend
Definition Agent.h:158
std::uint64_t trainingSamples
Definition Agent.h:156
std::uint64_t episodes
Definition Agent.h:153
std::uint64_t steps
Definition Agent.h:154
In-memory versioned replay evidence, independent of learned model state.
Definition Agent.h:89
Observation initial
Definition Agent.h:94
std::vector< TraceStep > steps
Definition Agent.h:95
std::string schemaId
Definition Agent.h:90
std::uint32_t schemaVersion
Definition Agent.h:91
std::uint64_t environmentSeed
Definition Agent.h:92