DecisionTransformer-Unity-Sim / Upload /01_Source_Code /Unity_Full50_Code_Snapshot /EvaluationProtocol.cs
Download Upload/01_Source_Code/Unity_Full50_Code_Snapshot/EvaluationProtocol.cs from code3939/DecisionTransformer-Unity-Sim: direct link, hf CLI and curl.
- Browser
- Download file 8.02 kB
-
https://huggingface.co/code3939/DecisionTransformer-Unity-Sim/resolve/main/Upload/01_Source_Code/Unity_Full50_Code_Snapshot/EvaluationProtocol.cs
- Command line
-
hf download hf://code3939/DecisionTransformer-Unity-Sim/Upload/01_Source_Code/Unity_Full50_Code_Snapshot/EvaluationProtocol.cs
-
curl -L -o EvaluationProtocol.cs https://huggingface.co/code3939/DecisionTransformer-Unity-Sim/resolve/main/Upload/01_Source_Code/Unity_Full50_Code_Snapshot/EvaluationProtocol.cs
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); | |
| } | |
| [] | |
| 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>(); | |
| } | |
| [] | |
| public class InputTrace | |
| { | |
| public int episode, timestep, valid_length; | |
| public float[] observations, actions, returns_to_go, selected_action; | |
| public int[] timesteps; | |
| } | |
| [] | |
| 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); | |
| } | |
| } | |