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