Repository navigation
Expand file tree
/
Copy pathtrainingData.py
More file actions
137 lines (123 loc) · 4.84 KB
/
Copy pathtrainingData.py
File metadata and controls
137 lines (123 loc) · 4.84 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
# Copyright 2017 Martin Haesemeyer. All rights reserved.
#
# Licensed under the MIT license
"""
Module to generate, save and load training and test data and provide access to that data in convenient batches
"""
import numpy as np
from core import GradientData
from zf_simulators import TrainingSimulation
class CircGradientTrainer(TrainingSimulation):
"""
Class to run a gradient simulation, creating training data in the process
The simulation will at each timepoint simulate what could have happened
500 ms into the future for each selectable behavior but only one path
will actually be chosen to advance the simulation to avoid massive branching
"""
def __init__(self, radius, t_min, t_max):
"""
Creates a new GradientSimulation
:param radius: The arena radius in mm
:param t_min: The center temperature
:param t_max: The edge temperature
"""
super().__init__()
self.radius = radius
self.t_min = t_min
self.t_max = t_max
def temperature(self, x, y, a=0):
"""
Returns the temperature at the given positions
"""
r = np.sqrt(x**2 + y**2) # this is a circular arena so compute radius
return (r / self.radius) * (self.t_max - self.t_min) + self.t_min
def out_of_bounds(self, x, y):
"""
Detects whether the given x-y position is out of the arena
:param x: The x position
:param y: The y position
:return: True if the given position is outside the arena, false otherwise
"""
# circular arena, compute radial position of point and compare to arena radius
r = np.sqrt(x**2 + y**2)
return r > self.radius
class LinGradientTrainer(TrainingSimulation):
"""
Class for generating training data in a linear gradient
"""
def __init__(self, xmax, ymax, t_min, t_max):
"""
Creates a new LinerGradient simulation
:param xmax: The maximal x-coordinate (gradient direction)
:param ymax: The maximal y-coordinate (neutral direction)
:param t_min: The minimum temperature at left edge
:param t_max: The maximal temperature at right edge
"""
super().__init__()
self.xmax = xmax
self.ymax = ymax
self.t_min = t_min
self.t_max = t_max
def temperature(self, x, y, a=0):
"""
Returns the temperature at the given positions
"""
return (x / self.xmax) * (self.t_max - self.t_min) + self.t_min
def out_of_bounds(self, x, y):
"""
Detects whether the given x-y position is out of the arena
:param x: The x position
:param y: The y position
:return: True if the given position is outside the arena, false otherwise
"""
if x < 0 or x > self.xmax:
return True
if y < 0 or y > self.ymax:
return True
return False
def run_simulation(self, nsteps):
"""
Forward run of random gradient exploration
:param nsteps: The number of steps to simulate
:return: The position and heading in the gradient at each timepoint
"""
spos = np.array([self.xmax // 2, self.ymax // 2, 0])
return self.sim_forward(nsteps, spos, "N").copy()
# The bout frequency to use during the virtual navigation for training data generation
TRAIN_BOUT_FREQ = 1
if __name__ == '__main__':
import matplotlib.pyplot as pl
import seaborn as sns
response = ""
while response not in ["y", "n"]:
response = input("Run simulation with default arena? [y/n]:")
if response == "y":
n_steps = int(input("Number of steps to perform?"))
gradsim = CircGradientTrainer(100, 22, 37)
gradsim.p_move *= TRAIN_BOUT_FREQ # Adjust bout frequency during training data navigation
print("Running radial simulation, inside-out")
pos = gradsim.run_simulation(n_steps)
pl.figure()
pl.plot(pos[:, 0], pos[:, 1])
pl.xlabel("X position [mm]")
pl.ylabel("Y position [mm]")
sns.despine()
print("Generating gradient data")
grad_data = gradsim.create_dataset(pos)
all_in = grad_data.model_in_raw
all_out = grad_data.model_out_raw
print("Running radial simulation, outside-in")
gradsim = CircGradientTrainer(100, 37, 22)
gradsim.p_move *= TRAIN_BOUT_FREQ
pos = gradsim.run_simulation(n_steps)
pl.figure()
pl.plot(pos[:, 0], pos[:, 1])
pl.xlabel("X position [mm]")
pl.ylabel("Y position [mm]")
sns.despine()
print("Generating gradient data")
grad_data = gradsim.create_dataset(pos)
all_in = np.r_[all_in, grad_data.model_in_raw]
all_out = np.r_[all_out, grad_data.model_out_raw]
grad_data = GradientData(all_in, all_out, grad_data.pred_window)
print("Done")