File size: 15,806 Bytes
b7e9b58 | 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 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 | #if UNITY_EDITOR
using System;
using System.Collections.Generic;
using System.IO;
using System.Linq;
using Unity.InferenceEngine;
using Unity.MLAgents;
using Unity.MLAgents.Policies;
using UnityEditor;
using UnityEditor.SceneManagement;
using UnityEngine;
using UnityEngine.SceneManagement;
// Batch entry points run only when explicitly requested with -executeMethod.
[InitializeOnLoad]
public static class RevisionEvaluationBatch
{
const string Pending = "RevisionEvaluationBatch.pending";
const string Failure = "RevisionEvaluationBatch.failure";
const string Started = "RevisionEvaluationBatch.started";
const string Played = "RevisionEvaluationBatch.played";
const string Done = "RevisionEvaluationBatch.done";
const string PpoModel = "Assets/Model/V12/V12 PPO.onnx";
const string DtModel = "Assets/Model/FinalModel/E_1_DT_C_5.onnx";
[Serializable] public class Job
{
public string id, policy, scene, output, report;
public string[] models;
public int episodes = 2, targets = 10, seed = 42, max_steps = 1000;
public float initial_rtg = 35f;
public double timeout_seconds = 7200;
}
[Serializable] class Verification
{
public string status, error, unity_version, gpu, graphics_api, utc;
public bool compilation_passed, reset_tests_passed;
public EvaluationModelParityTests.Result parity;
}
[Serializable] class RunReport
{
public string status, error, utc, unity_version, gpu, graphics_api;
public Job job;
public string source_scene = "Assets/Scenes/Shooting.unity";
public string rig_source, ppo_model, ppo_inference_device;
public bool ppo_deterministic;
public int[] ppo_discrete_branch_sizes;
public float fixed_delta_time, time_scale, rotation_speed, shoot_distance, step_penalty;
public string[] files;
public double wall_seconds;
}
static RevisionEvaluationBatch()
{
EditorApplication.update += Monitor;
Application.logMessageReceived += OnLog;
}
public static string Argument(string name)
{
var args = Environment.GetCommandLineArgs();
for (int i = 0; i + 1 < args.Length; i++) if (args[i] == name) return args[i + 1];
throw new ArgumentException("Missing command-line argument " + name);
}
public static void Verify()
{
var result = new Verification {
status = "failed", unity_version = Application.unityVersion,
gpu = SystemInfo.graphicsDeviceName, graphics_api = SystemInfo.graphicsDeviceType.ToString(),
utc = DateTime.UtcNow.ToString("o"), compilation_passed = true
};
int code = 1;
try
{
if (SystemInfo.graphicsDeviceType == UnityEngine.Rendering.GraphicsDeviceType.Null)
throw new InvalidOperationException("GPU verification requires a graphics device; omit -nographics.");
EvaluationSmokeTests.Run();
result.reset_tests_passed = true;
result.parity = EvaluationModelParityTests.Run(
AssetDatabase.LoadAssetAtPath<ModelAsset>(DtModel), Argument("-evaluationFixtures"));
result.status = "passed";
code = 0;
}
catch (Exception error) { result.error = error.ToString(); Debug.LogException(error); }
finally
{
WriteJson(Argument("-evaluationReport"), result);
EditorApplication.Exit(code);
}
}
public static void Run()
{
try
{
if (!Application.isBatchMode) throw new InvalidOperationException("Use this entry point only in batch mode on the working copy.");
Job job = JsonUtility.FromJson<Job>(File.ReadAllText(Argument("-evaluationJob")));
if (job == null || job.episodes < 1 || (job.targets != 10 && job.targets != 15 && job.targets != 20) ||
(job.policy != "PPO" && job.policy != "DT") || job.models == null || job.models.Length == 0)
throw new InvalidOperationException("Invalid evaluation job.");
if (Directory.Exists(job.output) && Directory.GetFiles(job.output, "*.json", SearchOption.AllDirectories).Length > 0)
throw new InvalidOperationException("Use a fresh result directory for every job.");
Directory.CreateDirectory(job.output);
var report = Prepare(job);
WriteJson(job.report, report);
SessionState.SetString(Pending, JsonUtility.ToJson(job));
SessionState.SetString(Failure, "");
SessionState.SetString(Started, DateTime.UtcNow.ToString("o"));
SessionState.SetBool(Played, false);
SessionState.SetBool(Done, false);
EditorApplication.isPlaying = true;
}
catch (Exception error)
{
Debug.LogException(error);
EditorApplication.Exit(1);
}
}
static RunReport Prepare(Job job)
{
var scene = EditorSceneManager.OpenScene("Assets/Scenes/Shooting.unity", OpenSceneMode.Single);
var ppo = scene.GetRootGameObjects().SelectMany(g => g.GetComponentsInChildren<DroneAgent_For_Testing>(true)).Single();
string rigSource = HierarchyPath(ppo.transform);
var barrel = ppo.gunBarrel;
var targets = new List<EnemyDummy>(ppo.enemies);
if (barrel == null || targets.Count != 10 || targets.Any(e => e == null) || targets.Distinct().Count() != 10)
throw new InvalidOperationException("Expected the original PPO rig with ten unique targets.");
var behavior = ppo.GetComponent<BehaviorParameters>();
var requester = ppo.GetComponent<DecisionRequester>();
if (behavior == null || requester == null || AssetDatabase.GetAssetPath(behavior.Model) != PpoModel)
throw new InvalidOperationException("Original PPO rig/model differs from the inspected project.");
if (behavior.BrainParameters.VectorObservationSize != 9 ||
behavior.BrainParameters.ActionSpec.NumContinuousActions != 2 ||
!behavior.BrainParameters.ActionSpec.BranchSizes.SequenceEqual(new[] { 2, 2 }))
throw new InvalidOperationException("PPO observation/action contract mismatch.");
// Disable the original scene before configuring anything; activate only the common rig.
foreach (var go in scene.GetRootGameObjects()) go.SetActive(false);
foreach (var enemy in scene.GetRootGameObjects().SelectMany(g => g.GetComponentsInChildren<EnemyDummy>(true)))
enemy.gameObject.SetActive(false);
foreach (var agent in scene.GetRootGameObjects().SelectMany(g => g.GetComponentsInChildren<Agent>(true)))
{
agent.enabled = false;
if (agent.gameObject != ppo.gameObject) agent.gameObject.SetActive(false);
}
foreach (var controller in scene.GetRootGameObjects().SelectMany(g => g.GetComponentsInChildren<DTController_For_Testing>(true)))
{
controller.enabled = false;
if (controller.gameObject != ppo.gameObject) controller.gameObject.SetActive(false);
}
foreach (var decision in scene.GetRootGameObjects().SelectMany(g => g.GetComponentsInChildren<DecisionRequester>(true)))
decision.enabled = false;
// 15/20-target conditions extend the same collider/motion settings, preserving the first ten and their order.
for (int i = targets.Count; i < job.targets; i++)
{
var clone = UnityEngine.Object.Instantiate(targets[i % 10].gameObject, targets[i % 10].transform.parent);
clone.name = "RevisionTarget_" + i;
clone.SetActive(false);
targets.Add(clone.GetComponent<EnemyDummy>());
}
foreach (var enemy in targets)
{
if (enemy.GetComponent<Collider>() == null || (ppo.enemyLayer.value & (1 << enemy.gameObject.layer)) == 0 ||
enemy.GetComponentsInChildren<Rigidbody>(true).Any(r => !r.isKinematic) ||
enemy.GetComponentsInChildren<Animator>(true).Any(a => a.enabled) ||
enemy.GetComponents<MonoBehaviour>().Any(b => b != enemy && b.enabled))
throw new InvalidOperationException("Target has unexpected collision or external movement settings: " + enemy.name);
ActivateAncestors(enemy.transform.parent);
}
ActivateAncestors(ppo.transform);
foreach (var target in targets) target.gameObject.SetActive(false); // ResetEnvironment owns first activation.
foreach (var other in ppo.GetComponents<Agent>().Where(a => a != ppo).ToArray())
UnityEngine.Object.DestroyImmediate(other);
ppo.maxTestEpisodes = job.episodes;
ppo.maxEpisodeSteps = job.max_steps;
ppo.testSeed = job.seed;
ppo.EnemyCount = job.targets;
ppo.enemies = targets;
ppo.rotationSpeed = 100f;
ppo.shootDistance = 50f;
ppo.stepPenalty = EvaluationReward.DefaultStepPenalty;
ppo.outputFileName = "V12_PPO.json";
behavior.BehaviorType = BehaviorType.InferenceOnly;
requester.DecisionPeriod = 1;
requester.DecisionStep = 0;
requester.TakeActionsBetweenDecisions = false;
if (job.policy == "PPO")
{
if (job.models.Length != 1 || job.models[0] != PpoModel) throw new InvalidOperationException("Unexpected PPO model.");
ppo.enabled = true;
requester.enabled = true;
}
else
{
// No PPO component is allowed to initialize an Academy/policy during DT evaluation.
var go = ppo.gameObject;
var mask = ppo.enemyLayer;
UnityEngine.Object.DestroyImmediate(requester);
UnityEngine.Object.DestroyImmediate(ppo);
var controller = go.GetComponent<DTController_For_Testing>() ?? go.AddComponent<DTController_For_Testing>();
controller.modelAssets = job.models.Select(path => AssetDatabase.LoadAssetAtPath<ModelAsset>(path)).ToList();
if (controller.modelAssets.Any(model => model == null)) throw new InvalidOperationException("A requested model is missing.");
controller.gunBarrel = barrel;
controller.enemies = targets;
controller.enemyLayer = mask;
controller.EnemyCount = job.targets;
controller.maxTestEpisodes = job.episodes;
controller.maxEpisodeSteps = job.max_steps;
controller.testSeed = job.seed;
controller.initialRTG = job.initial_rtg;
controller.RTG_TEST = false;
controller.policyMode = DTController_For_Testing.PolicyMode.AutoFromModelName;
controller.rotationSpeed = 100f;
controller.shootDistance = 50f;
controller.stepPenalty = EvaluationReward.DefaultStepPenalty;
controller.traceEpisodes = 2;
controller.laserLine = controller.laserLineToEnemy = null;
controller.enabled = true;
}
// Keep the original time scale and fixed timestep; target simulation advances by action count.
var activeAgents = scene.GetRootGameObjects().SelectMany(g => g.GetComponentsInChildren<Agent>(true))
.Where(a => a.isActiveAndEnabled).ToArray();
var activeDt = scene.GetRootGameObjects().SelectMany(g => g.GetComponentsInChildren<DTController_For_Testing>(true))
.Where(a => a.isActiveAndEnabled).ToArray();
if (activeAgents.Length != (job.policy == "PPO" ? 1 : 0) || activeDt.Length != (job.policy == "DT" ? 1 : 0))
throw new InvalidOperationException("The scene must have exactly one active evaluator.");
Directory.CreateDirectory(Path.GetDirectoryName(job.scene));
if (!EditorSceneManager.SaveScene(scene, job.scene)) throw new IOException("Could not save evaluation scene.");
AssetDatabase.SaveAssets();
return new RunReport {
status = "running", job = job, utc = DateTime.UtcNow.ToString("o"),
unity_version = Application.unityVersion, gpu = SystemInfo.graphicsDeviceName,
graphics_api = SystemInfo.graphicsDeviceType.ToString(), rig_source = rigSource,
ppo_model = AssetDatabase.GetAssetPath(behavior.Model),
ppo_inference_device = behavior.InferenceDevice.ToString(), ppo_deterministic = behavior.DeterministicInference,
ppo_discrete_branch_sizes = behavior.BrainParameters.ActionSpec.BranchSizes,
fixed_delta_time = Time.fixedDeltaTime, time_scale = Time.timeScale,
rotation_speed = 100f, shoot_distance = 50f, step_penalty = EvaluationReward.DefaultStepPenalty
};
}
static void ActivateAncestors(Transform transform)
{
if (transform == null) return;
ActivateAncestors(transform.parent);
transform.gameObject.SetActive(true);
}
static string HierarchyPath(Transform transform)
{
return transform.parent == null ? transform.name : HierarchyPath(transform.parent) + "/" + transform.name;
}
static void OnLog(string condition, string stack, LogType type)
{
if (string.IsNullOrEmpty(SessionState.GetString(Pending, ""))) return;
if (type == LogType.Error || type == LogType.Exception || type == LogType.Assert)
SessionState.SetString(Failure, condition + "\n" + stack);
}
static void Monitor()
{
string json = SessionState.GetString(Pending, "");
if (string.IsNullOrEmpty(json)) return;
try
{
var job = JsonUtility.FromJson<Job>(json);
if (EditorApplication.isPlaying) SessionState.SetBool(Played, true);
double elapsed = (DateTime.UtcNow - DateTime.Parse(SessionState.GetString(Started, ""), null,
System.Globalization.DateTimeStyles.RoundtripKind)).TotalSeconds;
var files = Directory.GetFiles(job.output, "*.json", SearchOption.AllDirectories);
string failure = SessionState.GetString(Failure, "");
if (elapsed > job.timeout_seconds) failure = "Evaluation exceeded the job timeout.";
bool finished = SessionState.GetBool(Played, false) && files.Length == job.models.Length;
if (SessionState.GetBool(Played, false) && !EditorApplication.isPlayingOrWillChangePlaymode && !SessionState.GetBool(Done, false))
failure = "Play mode stopped before expected results were saved.";
if (!finished && string.IsNullOrEmpty(failure)) return;
SessionState.SetBool(Done, true);
if (EditorApplication.isPlayingOrWillChangePlaymode)
{
if (!string.IsNullOrEmpty(failure)) SessionState.SetString(Failure, failure);
EditorApplication.isPlaying = false;
return;
}
var report = JsonUtility.FromJson<RunReport>(File.ReadAllText(job.report));
report.status = string.IsNullOrEmpty(failure) && finished ? "passed" : "failed";
report.error = failure;
report.files = files;
report.wall_seconds = elapsed;
WriteJson(job.report, report);
SessionState.EraseString(Pending);
EditorApplication.Exit(report.status == "passed" ? 0 : 1);
}
catch (Exception error)
{
SessionState.EraseString(Pending);
Debug.LogException(error);
EditorApplication.Exit(1);
}
}
static void WriteJson(string path, object value)
{
Directory.CreateDirectory(Path.GetDirectoryName(Path.GetFullPath(path)));
File.WriteAllText(path, JsonUtility.ToJson(value, true));
}
}
#endif
|