forked from SjulsonLab/RPi4_behavior_boxes
-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathself_admin_task.py
More file actions
308 lines (268 loc) · 11.5 KB
/
Copy pathself_admin_task.py
File metadata and controls
308 lines (268 loc) · 11.5 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
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
# python3: headfixed_task.py
"""
author: tian qiu
date: 2023-01-26
name: self_admin_Task.py
goal: self administration task adpated via Mitch's instruction
description:
an updated test version of self_admin_Task.py
"""
import importlib
from transitions import Machine
from transitions import State
from transitions.extensions.states import add_state_features, Timeout
import pysistence, collections
from icecream import ic
import logging
import time
from datetime import datetime
import os
from gpiozero import PWMLED, LED, Button
from colorama import Fore, Style
import logging.config
from time import sleep
import random
import threading
import matplotlib
import matplotlib.pyplot as plt
import matplotlib.figure as fg
import numpy as np
logging.config.dictConfig(
{
"version": 1,
"disable_existing_loggers": True,
}
)
# all modules above this line will have logging disabled
import behavbox
# adding timing capability to the state machine
@add_state_features(Timeout)
class TimedStateMachine(Machine):
pass
class SelfAdminTask(object):
# Define states. States where the animals is waited to make their decision
def __init__(self, **kwargs): # name and session_info should be provided as kwargs
# if no name or session, make fake ones (for testing purposes)
if kwargs.get("name", None) is None:
self.name = "name"
print(
Fore.RED
+ Style.BRIGHT
+ "Warning: no name supplied; making fake one"
+ Style.RESET_ALL
)
else:
self.name = kwargs.get("name", None)
if kwargs.get("session_info", None) is None:
print(
Fore.RED
+ Style.BRIGHT
+ "Warning: no session_info supplied; making fake one"
+ Style.RESET_ALL
)
from fake_session_info import fake_session_info
self.session_info = fake_session_info
else:
self.session_info = kwargs.get("session_info", None)
ic(self.session_info)
# initialize the state machine
self.states = [
State(name='standby',
on_enter=["enter_standby"],
on_exit=["exit_standby"]),
Timeout(name='reward_available',
on_enter=["enter_reward_available"],
on_exit=["exit_reward_available"],
timeout=self.session_info["reward_timeout"],
on_timeout=["restart"])
]
self.transitions = [
['start_trial', 'standby', 'reward_available'], # format: ['trigger', 'origin', 'destination']
['restart', 'reward_available', 'standby']
]
self.machine = TimedStateMachine(
model=self,
states=self.states,
transitions=self.transitions,
initial='standby'
)
self.trial_running = False
# trial statistics
self.innocent = True
self.trial_number = 0
self.error_count = 0
self.error_list = []
self.error_repeat = False
self.lever_pressed_time = 0.0
self.lever_press_interval = self.session_info["lever_press_interval"]
# self.reward_time_start = None # for reward_available state time keeping purpose
self.reward_time = 10 # sec. could be incorporate into the session_info; available time for reward
self.reward_times_up = False
self.reward_pump = self.session_info["reward_pump"]
self.reward_size = self.session_info["reward_size"]
self.active_press = 0
self.inactive_press = 0
self.timeline_active_press = []
self.active_press_count_list = []
self.timeline_inactive_press = []
self.inactive_press_count_list = []
self.left_poke_count = 0
self.right_poke_count = 0
self.timeline_left_poke = []
self.left_poke_count_list = []
self.timeline_right_poke = []
self.right_poke_count_list = []
self.event_name = ""
# initialize behavior box
self.box = behavbox.BehavBox(self.session_info)
self.pump = self.box.pump
self.treadmill = self.box.treadmill
# self.distance_initiation = self.session_info['treadmill_setup']['distance_initiation']
# self.distance_cue = self.session_info['treadmill_setup']['distance_cue']
# self.distance_buffer = None
# self.distance_diff = 0
# for refining the lick detection
self.lick_count = 0
self.side_mice_buffer = None
self.LED_blink = False
try:
self.lick_threshold = self.session_info["lick_threshold"]
except:
print("No lick_threshold defined in session_info. Therefore, default defined as 2 \n")
self.lick_threshold = 1
# session_statistics
self.total_reward = 0
########################################################################
# functions called when state transitions occur
########################################################################
def run(self):
if self.box.event_list:
self.event_name = self.box.event_list.popleft()
else:
self.event_name = ""
if self.state == "standby":
pass
elif self.state == "reward_available":
if self.event_name == "reserved_rx1_pressed":
lever_pressed_time_temp = time.time()
lever_pressed_dt = lever_pressed_time_temp - self.lever_pressed_time
if lever_pressed_dt >= self.lever_press_interval:
self.pump.reward(self.reward_pump, self.reward_size)
self.lever_pressed_time = lever_pressed_time_temp
self.total_reward += 1
# self.active_press += 1
# self.active_press_count_list.append(self.left_poke_count)
# self.timeline_active_press.append(time.time())
# elif self.event_name == "reserved_rx2_pressed":
#
# self.inactive_press += 1
# self.inactive_press_count_list.append(self.right_poke_count)
# self.timeline_inactive_press.append(time.time())
# Lick detection:
# if self.event_name == "left_IR_entry":
# self.left_poke_count += 1
# self.left_poke_count_list.append(self.left_poke_count)
# self.timeline_left_poke.append(time.time())
# self.lick_count += 1
# look for keystrokes
self.box.check_keybd()
def enter_standby(self):
logging.info(";" + str(time.time()) + ";[transition];enter_standby;" + str(self.error_repeat))
# self.cue_off('all')
self.update_plot_choice()
# self.update_plot_error()
self.trial_running = False
# self.reward_error = False
# if self.early_lick_error:
# self.error_list.append("early_lick_error")
# self.early_lick_error = False
# self.lick_count = 0
# self.side_mice_buffer = None
print(str(time.time()) + ", Total reward up till current session: " + str(self.total_reward))
logging.info(";" + str(time.time()) + ";[trial];trial_" + str(self.trial_number) + ";" + str(self.error_repeat))
def exit_standby(self):
# self.error_repeat = False
logging.info(";" + str(time.time()) + ";[transition];exit_standby;" + str(self.error_repeat))
self.box.event_list.clear()
pass
def enter_reward_available(self):
logging.info(";" + str(time.time()) + ";[transition];enter_reward_available;" + str(self.error_repeat))
self.trial_running = True
# self.reward_times_up = False
def exit_reward_available(self):
logging.info(";" + str(time.time()) + ";[transition];exit_reward_available;" + str(self.error_repeat))
# self.cue_off('sound2')
# self.reward_times_up = True
self.pump.reward("vaccum", 0)
# if self.multiple_choice_error:
# logging.info(";" + str(time.time()) + ";[error];multiple_choice_error;" + str(self.error_repeat))
# self.error_repeat = False
# self.error_list.append('multiple_choice_error')
# self.multiple_choice_error = False
# elif self.lick_count == 0:
# logging.info(";" + str(time.time()) + ";[error];no_choice_error;" + str(self.error_repeat))
# self.error_repeat = True
# self.error_list.append('no_choice_error')
# self.lick_count = 0
# self.reward_time_start = None
def update_plot(self):
fig, axes = plt.subplots(1, 1, )
axes.plot([1, 2], [1, 2], color='green', label='test')
self.box.check_plot(fig)
def update_plot_error(self):
error_event = self.error_list
labels, counts = np.unique(error_event, return_counts=True)
ticks = range(len(counts))
fig, ax = plt.subplots(1, 1, )
ax.bar(ticks, counts, align='center', tick_label=labels)
# plt.xticks(ticks, labels)
# plt.title(session_name)
ax = plt.gca()
ax.set_xticks(ticks, labels)
ax.set_xticklabels(labels=labels, rotation=70)
self.box.check_plot(fig)
def update_plot_choice(self, save_fig=False):
trajectory_active = self.left_poke_count_list
time_active = self.timeline_left_poke
trajectory_inactive = self.right_poke_count_list
time_inactive = self.timeline_right_poke
fig, ax = plt.subplots(1, 1, )
print(type(fig))
ax.plot(time_active, trajectory_active, color='b', marker="o", label='active_trajectory')
ax.plot(time_inactive, trajectory_inactive, color='r', marker="o", label='inactive_trajectory')
if save_fig:
plt.savefig(self.session_info['basedir'] + "/" + self.session_info['basename'] + "/" + \
self.session_info['basename'] + "_lever_choice_plot" + '.png')
self.box.check_plot(fig)
def integrate_plot(self, save_fig=False):
fig, ax = plt.subplots(2, 1)
trajectory_left = self.active_press
time_active_press = self.timeline_active_press
trajectory_right = self.right_poke_count_list
time_inactive_press = self.timeline_inactive_press
print(type(fig))
ax[0].plot(time_active_press, trajectory_left, color='b', marker="o", label='left_lick_trajectory')
ax[0].plot(time_inactive_press, trajectory_right, color='r', marker="o", label='right_lick_trajectory')
error_event = self.error_list
labels, counts = np.unique(error_event, return_counts=True)
ticks = range(len(counts))
ax[1].bar(ticks, counts, align='center', tick_label=labels)
# plt.xticks(ticks, labels)
# plt.title(session_name)
ax[1] = plt.gca()
ax[1].set_xticks(ticks, labels)
ax[1].set_xticklabels(labels=labels, rotation=70)
if save_fig:
plt.savefig(self.session_info['basedir'] + "/" + self.session_info['basename'] + "/" + \
self.session_info['basename'] + "_summery" + '.png')
self.box.check_plot(fig)
########################################################################
# methods to start and end the behavioral session
########################################################################
def start_session(self):
ic("TODO: start video")
self.box.video_start()
def end_session(self):
ic("TODO: stop video")
self.update_plot_choice(save_fig=True)
self.box.video_stop()