code3939's picture
Upload 765 files
b7e9b58 verified
Raw History Blame Contribute Delete
8.02 kB
using System;
using System.Collections.Generic;
using System.IO;
using System.Security.Cryptography;
using System.Text;
using UnityEngine;
// Shared by PPO, DT and BC evaluation. No Unity editor APIs required.
public static class EvaluationProtocol
{
public const string Version = "aligned-evaluation-v2.1";
public static int EpisodeSeed(int masterSeed, int zeroBasedEpisode)
{
if (zeroBasedEpisode < 0) throw new ArgumentOutOfRangeException(nameof(zeroBasedEpisode));
return checked(masterSeed + zeroBasedEpisode);
}
[Serializable]
public class InitialState
{
public int episode;
public int seed;
public Vector3 agent_position;
public Quaternion agent_rotation;
public Vector3 agent_scale;
public Vector3 barrel_position;
public Vector3 barrel_scale;
public Quaternion barrel_rotation;
public List<EnemyDummy.EvaluationState> targets = new List<EnemyDummy.EvaluationState>();
}
[Serializable]
public class InputTrace
{
public int episode, timestep, valid_length;
public float[] observations, actions, returns_to_go, selected_action;
public int[] timesteps;
}
[Serializable]
public class Metadata
{
public string protocol_version = Version;
public string run_id;
public string model_name;
public string policy_type;
public int target_count;
public int master_seed;
public int episodes_requested;
public int max_action_steps;
public float fixed_delta_time;
public float rotation_speed;
public float shoot_distance;
public float step_penalty;
public int enemy_layer_mask;
public float initial_rtg;
public bool bc_zero_rtg;
public string sequence_protocol = "dataset-shifted-action-right-pad32-v2";
public string accuracy_protocol = "hits-per-fire-request-percent-v1";
public string action_history_protocol = "raw-continuous-and-actual-fire-v1";
public string inference_backend;
public int[] ppo_discrete_branch_sizes;
public bool ppo_deterministic_inference;
public bool queries_hit_triggers;
public bool rtg_sensitivity;
public string evaluated_utc = DateTime.UtcNow.ToString("o");
public List<float> episode_terminal_rtg = new List<float>();
public List<InputTrace> input_traces = new List<InputTrace>();
public string reward_protocol;
public string unity_version;
public List<int> episode_seeds = new List<int>();
public List<string> initial_state_sha256 = new List<string>();
public List<InitialState> initial_states = new List<InitialState>();
public List<float> max_abs_rtg_input = new List<float>();
public List<int> episode_action_steps = new List<int>();
public List<int> episode_shots_fired = new List<int>();
public List<int> episode_shots_hit = new List<int>();
public List<int> episode_raycasts = new List<int>();
}
public static void ValidateConfiguration(float rotationSpeed, float shootDistance,
float stepPenalty, float fixedDeltaTime)
{
if (!Finite(rotationSpeed) || rotationSpeed <= 0 || !Finite(shootDistance) || shootDistance <= 0 ||
!Finite(stepPenalty) || !Finite(fixedDeltaTime) || fixedDeltaTime <= 0)
throw new InvalidOperationException("Non-finite/invalid movement, reward or timestep settings.");
}
public static bool Finite(float value) { return !float.IsNaN(value) && !float.IsInfinity(value); }
public static InitialState ResetEnvironment(Transform root, Transform barrel,
List<EnemyDummy> enemies, int count, int masterSeed, int episode)
{
if (root == null || barrel == null || enemies == null || count <= 0 || count > enemies.Count)
throw new InvalidOperationException("Check root, barrel and EnemyCount/Enemies references.");
var unique = new HashSet<EnemyDummy>();
foreach (var e in enemies)
if (e == null || !unique.Add(e))
throw new InvalidOperationException("Enemies must contain unique, non-null entries in a fixed order.");
int seed = EpisodeSeed(masterSeed, episode);
// Lifecycle callbacks on deactivation must not consume our seeded initialization stream.
foreach (var e in enemies) e.gameObject.SetActive(false);
var saved = UnityEngine.Random.state;
var state = new InitialState {
episode = episode + 1, seed = seed, agent_position = root.position,
agent_rotation = root.rotation, barrel_rotation = barrel.localRotation,
agent_scale = root.lossyScale, barrel_position = barrel.position, barrel_scale = barrel.lossyScale
};
try
{
UnityEngine.Random.InitState(seed);
// Activate only after every target's random parameters have been drawn.
for (int i = 0; i < count; i++)
{
var e = enemies[i];
Vector3 pos = root.position + new Vector3(UnityEngine.Random.Range(-7f, 7f),
UnityEngine.Random.Range(2f, 3.5f), UnityEngine.Random.Range(-7f, 7f));
e.ResetForEvaluation(i, root, pos);
}
for (int i = 0; i < count; i++) enemies[i].gameObject.SetActive(true);
for (int i = 0; i < count; i++) state.targets.Add(enemies[i].CaptureEvaluationState());
}
finally { UnityEngine.Random.state = saved; }
Physics.SyncTransforms();
return state;
}
public static void RecordInitialState(Metadata metadata, InitialState state)
{
metadata.episode_seeds.Add(state.seed);
metadata.initial_states.Add(state);
// Episode/seed intentionally included; compare corresponding episodes only.
using (var sha = SHA256.Create())
{
byte[] bytes = sha.ComputeHash(Encoding.UTF8.GetBytes(JsonUtility.ToJson(state)));
metadata.initial_state_sha256.Add(BitConverter.ToString(bytes).Replace("-", "").ToLowerInvariant());
}
}
public static void AdvanceTargets(List<EnemyDummy> enemies, int completedActionSteps)
{
foreach (var e in enemies)
if (e != null && e.gameObject.activeSelf)
e.AdvanceEvaluation(completedActionSteps, Time.fixedDeltaTime);
Physics.SyncTransforms();
}
public static void StopTargets(List<EnemyDummy> enemies)
{
if (enemies == null) return;
foreach (var e in enemies) if (e != null) e.gameObject.SetActive(false);
}
public static string NewRunId()
{
return DateTime.UtcNow.ToString("yyyyMMdd_HHmmss_fff") + "_" + Guid.NewGuid().ToString("N").Substring(0, 8);
}
public static string OutputDirectory(string runId, int targetCount, bool rtgTest)
{
string outputRoot = Path.Combine(Application.persistentDataPath, "RevisionEvaluation");
// Explicit batch jobs can keep all artifacts inside their isolated workspace.
var arguments = Environment.GetCommandLineArgs();
for (int i = 0; i + 1 < arguments.Length; i++)
if (arguments[i] == "-evaluationOutput") outputRoot = Path.GetFullPath(arguments[i + 1]);
string path = Path.Combine(outputRoot, runId,
rtgTest ? "RTG_Test" : "EnemyCount" + targetCount);
Directory.CreateDirectory(path);
return path;
}
public static string SafeName(string name)
{
foreach (char c in Path.GetInvalidFileNameChars()) name = name.Replace(c, '_');
return name.Replace('/', '_').Replace('\\', '_');
}
public static void WriteNew(string path, string json)
{
using (var stream = new FileStream(path, FileMode.CreateNew, FileAccess.Write))
using (var writer = new StreamWriter(stream, new UTF8Encoding(false))) writer.Write(json);
}
}