Add CenterOfGravity detection / Add reward function

This commit is contained in:
Tobi
2019-05-17 18:23:24 +02:00
parent 7317135cd3
commit a7dd5527bb
11 changed files with 1152 additions and 44 deletions
@@ -4,9 +4,10 @@ using UnityEngine;
using OpenCvSharp;
using System.Threading.Tasks;
public class OnRenderImageTest : MonoBehaviour
public class ImageProcessor : MonoBehaviour
{
public MeshRenderer processedImageRenderer;
public Point CenterOfGravity { get; private set; }
// OpenCVSharp parameters
private Mat videoSourceImage;
@@ -41,6 +42,7 @@ public class OnRenderImageTest : MonoBehaviour
Texture2D tex = RenderTextureToTexture2D(source);
videoSourceImage = TextureToMat(tex);
cannyImage = ProcessImage(videoSourceImage);
CenterOfGravityTest(cannyImage);
processedTexture = MatToTexture(cannyImage);
Graphics.Blit(processedTexture, destination);
}
@@ -124,4 +126,18 @@ public class OnRenderImageTest : MonoBehaviour
texture.Apply();
return texture;
}
private void CenterOfGravityTest(Mat processedImage)
{
// find moments of the image
Moments m = new Moments(processedImage, true);
CenterOfGravity = new Point(m.M10 / m.M00, m.M01 / m.M00);
// show the image with a point mark at the centroid
Cv2.Circle(processedImage, CenterOfGravity, 5, new Scalar(128, 0, 0), -1);
Cv2.Flip(processedImage, processedImage, FlipMode.X);
Cv2.ImShow("Image with center", processedImage);
//Cv2.WaitKey(0);
}
}
+98 -2
View File
@@ -1,21 +1,117 @@
using MLAgents;
using OpenCvSharp;
using System;
using System.Collections;
using System.Collections.Generic;
using UnityEngine;
public class CozmoAgent : Agent
{
// Possible Actions
private const int STOP = 0;
private const int FORWARD = 1;
private const int RIGHT = 2;
private const int LEFT = 3;
// Used to determine different areas in the image (near to the center, far away)
private const float NEAR_AREA_PERCENTAGE_OFFSET = 0.2f;
private const float FAR_AREA_PERCENTAGE_OFFSET = 0.3f;
[Tooltip("The virtual Cozmo camera")]
public Camera renderCamera;
public Academy academy;
[Tooltip("Reference to the CozmoMovement script")]
public CozmoMovement movement;
public float timeBetweenDecisionsAtInference;
private float timeSinceDecision;
private Academy academy; // CozmoAcademy
private float timeSinceDecision; // time since last decision
private ImageProcessor onRenderImageTest; // reference to the ImageProcessor
private int nearAreaLimit = 0; // X coordinate limit for the near to the imagecenter area
private int farAreaLimit = 0; // X coordinate limit for the far away to the imagecenter area
// for testing
//private float[] floats = { 1.0f, 2.0f, 3.0f };
private void Start()
{
academy = FindObjectOfType(typeof(CozmoAcademy)) as CozmoAcademy;
onRenderImageTest = renderCamera.GetComponent<ImageProcessor>();
nearAreaLimit = (int)(renderCamera.targetTexture.width / 2 * NEAR_AREA_PERCENTAGE_OFFSET);
farAreaLimit = (int)(renderCamera.targetTexture.width / 2 * FAR_AREA_PERCENTAGE_OFFSET);
}
public void FixedUpdate()
{
WaitTimeInference();
// for testing
//AgentAction(floats, "ActionText");
}
// to be implemented by the developer
public override void AgentAction(float[] vectorAction, string textAction)
{
int action = Mathf.FloorToInt(vectorAction[0]);
Point centerOfGravity = onRenderImageTest.CenterOfGravity;
Vector3 targetPos = transform.position;
AddReward(-0.01f);
switch (action)
{
case STOP:
movement.Move(0);
break;
case FORWARD:
movement.Move(1);
break;
case RIGHT:
movement.Turn(1);
break;
case LEFT:
movement.Turn(-1);
break;
default:
movement.Move(0);
throw new ArgumentException("Invalid action value. Stop movement.");
}
// Render new image after movement in order to update the centerOfGravity
if (renderCamera != null)
{
renderCamera.Render();
}
// If centerOfGravity lies near to the center of the image horizontally
if (centerOfGravity.X > renderCamera.targetTexture.width / 2 - nearAreaLimit && centerOfGravity.X < renderCamera.targetTexture.width / 2 + nearAreaLimit)
{
Done();
SetReward(1);
print("Reward: +1");
}
else if (centerOfGravity.X > renderCamera.targetTexture.width / 2 - farAreaLimit && centerOfGravity.X < renderCamera.targetTexture.width / 2 + farAreaLimit)
{
Done();
SetReward(-1);
print("Reward: -1");
}
else
{
Done();
SetReward(-2);
print("Reward: -2");
}
}
// to be implemented by the developer
public override void AgentReset()
{
academy.AcademyReset();
}
private void WaitTimeInference()
{
if (renderCamera != null)
+23
View File
@@ -0,0 +1,23 @@
using System.Collections;
using System.Collections.Generic;
using UnityEngine;
public class SceneHelper : MonoBehaviour
{
[Tooltip("All the objects in this list will be deactivated/activated when the game is running")]
public List<GameObject> toggleInPlayMode;
public void Awake()
{
ToggleObjectList();
}
private void ToggleObjectList()
{
foreach (GameObject go in toggleInPlayMode)
{
go.SetActive(!go.activeSelf);
}
}
}
+11
View File
@@ -0,0 +1,11 @@
fileFormatVersion: 2
guid: bc64682be4f86884ebe75488344a51b2
MonoImporter:
externalObjects: {}
serializedVersion: 2
defaultReferences: []
executionOrder: 0
icon: {instanceID: 0}
userData:
assetBundleName:
assetBundleVariant: