| using System; |
| using System.Collections.Generic; |
| using UnityEngine; |
| using Unity.InferenceEngine; |
| using System.IO; |
| using Unity.MLAgents; |
| using Unity.MLAgents.Policies; |
| #if UNITY_EDITOR |
| using UnityEditor; |
| #endif |
|
|
| namespace Unity.MLAgentsExamples |
| { |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| public class ModelOverrider : MonoBehaviour |
| { |
| HashSet<string> k_SupportedExtensions = new HashSet<string> { "nn", "onnx" }; |
| const string k_CommandLineModelOverrideDirectoryFlag = "--mlagents-override-model-directory"; |
| const string k_CommandLineModelOverrideExtensionFlag = "--mlagents-override-model-extension"; |
| const string k_CommandLineQuitAfterEpisodesFlag = "--mlagents-quit-after-episodes"; |
| const string k_CommandLineQuitAfterSeconds = "--mlagents-quit-after-seconds"; |
| const string k_CommandLineQuitOnLoadFailure = "--mlagents-quit-on-load-failure"; |
|
|
| |
| Agent m_Agent; |
|
|
| |
| |
| private bool m_HaveProcessedCommandLine; |
|
|
| string m_BehaviorNameOverrideDirectory; |
|
|
| private string m_OriginalBehaviorName; |
|
|
| private List<string> m_OverrideExtensions = new List<string>(); |
|
|
| |
| Dictionary<string, ModelAsset> m_CachedModels = new Dictionary<string, ModelAsset>(); |
|
|
| |
| |
| int m_MaxEpisodes; |
|
|
| |
| DateTime m_Deadline = DateTime.MaxValue; |
|
|
| int m_NumSteps; |
| int m_PreviousNumSteps; |
| int m_PreviousAgentCompletedEpisodes; |
|
|
| bool m_QuitOnLoadFailure; |
| [Tooltip("Debug values to be used in place of the command line for overriding models.")] |
| public string debugCommandLineOverride; |
|
|
| |
| |
| static int s_PreviousAgentCompletedEpisodes; |
| static int s_PreviousNumSteps; |
|
|
| int TotalCompletedEpisodes |
| { |
| get { return m_PreviousAgentCompletedEpisodes + (m_Agent == null ? 0 : m_Agent.CompletedEpisodes); } |
| } |
|
|
| int TotalNumSteps |
| { |
| get { return m_PreviousNumSteps + m_NumSteps; } |
| } |
|
|
| public bool HasOverrides |
| { |
| get |
| { |
| GetAssetPathFromCommandLine(); |
| return !string.IsNullOrEmpty(m_BehaviorNameOverrideDirectory); |
| } |
| } |
|
|
| |
| |
| |
| public string OriginalBehaviorName |
| { |
| get |
| { |
| if (string.IsNullOrEmpty(m_OriginalBehaviorName)) |
| { |
| var bp = m_Agent.GetComponent<BehaviorParameters>(); |
| m_OriginalBehaviorName = bp.BehaviorName; |
| } |
|
|
| return m_OriginalBehaviorName; |
| } |
| } |
|
|
| public static string GetOverrideBehaviorName(string originalBehaviorName) |
| { |
| return $"Override_{originalBehaviorName}"; |
| } |
|
|
| |
| |
| |
| |
| void GetAssetPathFromCommandLine() |
| { |
| if (m_HaveProcessedCommandLine) |
| { |
| return; |
| } |
|
|
| var maxEpisodes = 0; |
| var timeoutSeconds = 0; |
|
|
| string[] commandLineArgsOverride = null; |
| if (!string.IsNullOrEmpty(debugCommandLineOverride) && Application.isEditor) |
| { |
| commandLineArgsOverride = debugCommandLineOverride.Split(' '); |
| } |
|
|
| var args = commandLineArgsOverride ?? Environment.GetCommandLineArgs(); |
| for (var i = 0; i < args.Length; i++) |
| { |
| if (args[i] == k_CommandLineModelOverrideDirectoryFlag && i < args.Length - 1) |
| { |
| m_BehaviorNameOverrideDirectory = args[i + 1].Trim(); |
| } |
| else if (args[i] == k_CommandLineModelOverrideExtensionFlag && i < args.Length - 1) |
| { |
| var overrideExtension = args[i + 1].Trim().ToLower(); |
| var isKnownExtension = k_SupportedExtensions.Contains(overrideExtension); |
| if (!isKnownExtension) |
| { |
| Debug.LogError($"loading unsupported format: {overrideExtension}"); |
| Application.Quit(1); |
| #if UNITY_EDITOR |
| EditorApplication.isPlaying = false; |
| #endif |
| } |
|
|
| m_OverrideExtensions.Add(overrideExtension); |
| } |
| else if (args[i] == k_CommandLineQuitAfterEpisodesFlag && i < args.Length - 1) |
| { |
| Int32.TryParse(args[i + 1], out maxEpisodes); |
| } |
| else if (args[i] == k_CommandLineQuitAfterSeconds && i < args.Length - 1) |
| { |
| Int32.TryParse(args[i + 1], out timeoutSeconds); |
| } |
| else if (args[i] == k_CommandLineQuitOnLoadFailure) |
| { |
| m_QuitOnLoadFailure = true; |
| } |
| } |
|
|
| if (!string.IsNullOrEmpty(m_BehaviorNameOverrideDirectory)) |
| { |
| |
| m_MaxEpisodes = maxEpisodes > 0 ? maxEpisodes : 1; |
| Debug.Log($"setting m_MaxEpisodes to {maxEpisodes}"); |
| } |
|
|
| if (timeoutSeconds > 0) |
| { |
| m_Deadline = DateTime.Now + TimeSpan.FromSeconds(timeoutSeconds); |
| Debug.Log($"setting deadline to {timeoutSeconds} from now."); |
| } |
|
|
| m_HaveProcessedCommandLine = true; |
| } |
|
|
| void OnEnable() |
| { |
| |
| m_PreviousNumSteps = s_PreviousNumSteps; |
| m_PreviousAgentCompletedEpisodes = s_PreviousAgentCompletedEpisodes; |
|
|
| m_Agent = GetComponent<Agent>(); |
|
|
| GetAssetPathFromCommandLine(); |
| if (HasOverrides) |
| { |
| OverrideModel(); |
| } |
| } |
|
|
| void OnDisable() |
| { |
| |
| |
| |
| s_PreviousAgentCompletedEpisodes = Mathf.Max(s_PreviousAgentCompletedEpisodes, TotalCompletedEpisodes); |
| s_PreviousNumSteps = Mathf.Max(s_PreviousNumSteps, TotalNumSteps); |
| } |
|
|
| void FixedUpdate() |
| { |
| if (m_MaxEpisodes > 0) |
| { |
| |
| |
| |
| |
| if (TotalCompletedEpisodes >= m_MaxEpisodes && TotalNumSteps > m_MaxEpisodes * m_Agent.MaxStep) |
| { |
| Debug.Log($"ModelOverride reached {TotalCompletedEpisodes} episodes and {TotalNumSteps} steps. Exiting."); |
| Application.Quit(0); |
| #if UNITY_EDITOR |
| EditorApplication.isPlaying = false; |
| #endif |
| } |
| else if (DateTime.Now >= m_Deadline) |
| { |
| Debug.Log( |
| $"Deadline exceeded. " + |
| $"{TotalCompletedEpisodes}/{m_MaxEpisodes} episodes and " + |
| $"{TotalNumSteps}/{m_MaxEpisodes * m_Agent.MaxStep} steps completed. Exiting."); |
| Application.Quit(0); |
| #if UNITY_EDITOR |
| EditorApplication.isPlaying = false; |
| #endif |
| } |
| } |
|
|
| m_NumSteps++; |
| } |
|
|
| public ModelAsset GetModelForBehaviorName(string behaviorName) |
| { |
| if (m_CachedModels.ContainsKey(behaviorName)) |
| { |
| return m_CachedModels[behaviorName]; |
| } |
|
|
| if (string.IsNullOrEmpty(m_BehaviorNameOverrideDirectory)) |
| { |
| Debug.Log($"No override directory set."); |
| return null; |
| } |
|
|
| |
| var overrideExtensions = (m_OverrideExtensions.Count > 0) |
| ? m_OverrideExtensions.ToArray() |
| : new[] { "nn", "onnx" }; |
|
|
| byte[] rawModel = null; |
| bool isOnnx = false; |
| string assetName = null; |
| foreach (var overrideExtension in overrideExtensions) |
| { |
| var assetPath = Path.Combine(m_BehaviorNameOverrideDirectory, $"{behaviorName}.{overrideExtension}"); |
| try |
| { |
| rawModel = File.ReadAllBytes(assetPath); |
| isOnnx = overrideExtension.Equals("onnx"); |
| assetName = "Override - " + Path.GetFileName(assetPath); |
| break; |
| } |
| catch (IOException) |
| { |
| |
| } |
| } |
|
|
| if (rawModel == null) |
| { |
| Debug.Log($"Couldn't load model file(s) for {behaviorName} in {m_BehaviorNameOverrideDirectory} (full path: {Path.GetFullPath(m_BehaviorNameOverrideDirectory)}"); |
|
|
| |
| m_CachedModels[behaviorName] = null; |
| return null; |
| } |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| var asset = LoadSentisModel(rawModel); |
| asset.name = assetName; |
| m_CachedModels[behaviorName] = asset; |
| return asset; |
| } |
|
|
| ModelAsset LoadSentisModel(byte[] rawModel) |
| { |
| var asset = ScriptableObject.CreateInstance<ModelAsset>(); |
| |
| |
| return asset; |
| } |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| |
| |
| |
| void OverrideModel() |
| { |
| bool overrideOk = false; |
| string overrideError = null; |
|
|
| m_Agent.LazyInitialize(); |
|
|
| ModelAsset ModelAsset = null; |
| try |
| { |
| ModelAsset = GetModelForBehaviorName(OriginalBehaviorName); |
| } |
| catch (Exception e) |
| { |
| overrideError = $"Exception calling GetModelForBehaviorName: {e}"; |
| } |
|
|
| if (ModelAsset == null) |
| { |
| if (string.IsNullOrEmpty(overrideError)) |
| { |
| overrideError = |
| $"Didn't find a model for behaviorName {OriginalBehaviorName}. Make " + |
| "sure the behaviorName is set correctly in the commandline " + |
| "and that the model file exists"; |
| } |
| } |
| else |
| { |
| var modelName = ModelAsset != null ? ModelAsset.name : "<null>"; |
| Debug.Log($"Overriding behavior {OriginalBehaviorName} for agent with model {modelName}"); |
| try |
| { |
| m_Agent.SetModel(GetOverrideBehaviorName(OriginalBehaviorName), ModelAsset); |
| overrideOk = true; |
| } |
| catch (Exception e) |
| { |
| overrideError = $"Exception calling Agent.SetModel: {e}"; |
| } |
| } |
|
|
| if (!overrideOk && m_QuitOnLoadFailure) |
| { |
| if (!string.IsNullOrEmpty(overrideError)) |
| { |
| Debug.LogWarning(overrideError); |
| } |
|
|
| Application.Quit(1); |
| #if UNITY_EDITOR |
| EditorApplication.isPlaying = false; |
| #endif |
| } |
| } |
| } |
| } |
|
|