code3939's picture
Upload 765 files
b7e9b58 verified
Raw History Blame Contribute Delete
15.8 kB
#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