using System; using System.Collections.Generic; using System.Linq; using Unity.InferenceEngine; using Unity.MLAgents.Inference.Utils; using Unity.MLAgents.Policies; namespace Unity.MLAgents.Inference { /// /// Tensor - A class to encapsulate a Tensor used for inference. /// /// This class contains the Array that holds the data array, the shapes, type and the /// placeholder in the execution graph. All the fields are editable in the inspector, /// allowing the user to specify everything but the data in a graphical way. /// [Serializable] internal class TensorProxy { public enum TensorType { Integer, FloatingPoint }; static readonly Dictionary k_TypeMap = new Dictionary() { { TensorType.FloatingPoint, typeof(float) }, { TensorType.Integer, typeof(int) } }; static readonly Dictionary k_DTypeMap = new Dictionary() { { TensorType.FloatingPoint, InferenceEngine.DataType.Float }, { TensorType.Integer, InferenceEngine.DataType.Int } }; public string name; public TensorType valueType; // Since Type is not serializable, we use the DisplayType for the Inspector public Type DataType => k_TypeMap[valueType]; public DataType DType => k_DTypeMap[valueType]; public int[] shape; [NonSerialized] public Tensor data; public BackendType Device => data.dataOnBackend.backendType; public long Height { get { return shape.Length >= 4 ? shape[^2] : 1; } } public long Width { get { return shape.Length >= 3 ? shape[^1] : 1; } } public long Channels { get { return shape.Length >= 4 ? shape[^3] : shape.Length == 3 ? shape[^2] : shape.Length == 2 ? shape[^1] : 1; } } ~TensorProxy() { Dispose(); } void Dispose() { if (data.dataOnBackend.backendType != BackendType.CPU) { data?.Dispose(); } } } internal static class TensorUtils { public static void ResizeTensor(TensorProxy tensor, int batch) { if (tensor.shape[0] == batch && tensor.data != null && tensor.data.Batch() == batch) { return; } tensor.data?.Dispose(); tensor.shape[0] = batch; var newTensorShape = new TensorShape(tensor.shape.Select(i => (int)i).ToArray()); tensor.data = CreateEmptyTensor(newTensorShape, tensor.DType); } public static Tensor CreateEmptyTensor(TensorShape shape, DataType dataType) { Tensor tensor = null; switch (dataType) { case DataType.Float: tensor = new Tensor(shape); break; case DataType.Int: tensor = new Tensor(shape); break; } return tensor; } internal static int[] TensorShapeFromSentis(TensorShape src) { if (src.rank == 2) { return new int[] { src.Batch(), src.Channels() }; } if (src.Height() == 1 && src.Width() == 1) { return new int[] { src.Batch(), src.Channels() }; } return new int[] { src.Batch(), src.Channels(), src.Height(), src.Width() }; } public static TensorProxy TensorProxyFromSentis(Tensor src, string nameOverride = null) { var shape = TensorShapeFromSentis(src.shape); return new TensorProxy { // name = nameOverride ?? src.name, name = nameOverride ?? "", valueType = src.dataType == DataType.Float ? TensorProxy.TensorType.FloatingPoint : TensorProxy.TensorType.Integer, shape = shape, data = src }; } /// /// Fill a specific batch of a TensorProxy with a given value /// /// /// The batch index to fill. /// public static void FillTensorBatch(TensorProxy tensorProxy, int batch, float fillValue) { var height = tensorProxy.data.Height(); var width = tensorProxy.data.Width(); var channels = tensorProxy.data.Channels(); tensorProxy.data.CompleteAllPendingOperations(); for (var h = 0; h < height; h++) { for (var w = 0; w < width; w++) { for (var c = 0; c < channels; c++) { ((Tensor)tensorProxy.data)[batch, c, h, w] = fillValue; } } } } /// /// Fill a pre-allocated Tensor with random numbers /// /// The pre-allocated Tensor to fill /// RandomNormal object used to populate tensor /// /// Throws when trying to fill a Tensor of type other than float /// /// /// Throws when the Tensor is not allocated /// public static void FillTensorWithRandomNormal( TensorProxy tensorProxy, RandomNormal randomNormal) { if (tensorProxy.DataType != typeof(float)) { throw new NotImplementedException("Only float data types are currently supported"); } if (tensorProxy.data == null) { throw new ArgumentNullException(); } tensorProxy.data.CompleteAllPendingOperations(); for (var i = 0; i < tensorProxy.data.Length(); i++) { ((Tensor)tensorProxy.data)[i] = (float)randomNormal.NextDouble(); } } } }