Repository navigation
Expand file tree
/
Copy pathData.py
More file actions
62 lines (53 loc) · 2.34 KB
/
Copy pathData.py
File metadata and controls
62 lines (53 loc) · 2.34 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
"""Handles data loading."""
import random
from libs.utils2 import opjD, lo
import libs.Segment_Data as Segment_Data
from Parameters import ARGS
class DataIndex(object):
"""
Index object, keeps track of position in data stack.
"""
def __init__(self, valid_data_moments, ctr, epoch_counter):
self.valid_data_moments = valid_data_moments
self.ctr = ctr
self.epoch_counter = epoch_counter
self.epoch_complete = False
class Data(object):
def get_segment_data(self):
"""Loads SegmentData from hdf5 segments"""
self.hdf5_runs_path = self.hdf5_segment_metadata_path = ARGS.data_path
self.hdf5_runs_path += '/hdf5/runs'
self.hdf5_segment_metadata_path += '/hdf5/segment_metadata'
Segment_Data.load_Segment_Data(self.hdf5_segment_metadata_path,
self.hdf5_runs_path)
def __init__(self):
self.get_segment_data()
# Load data indexes for training and validation
train_all_steer_path = ARGS.data_path + '/train_all_steer'
val_all_steer_path = ARGS.data_path + '/val_all_steer'
print('loading train_valid_data_moments...')
self.train_index = DataIndex(lo(train_all_steer_path), -1, 0)
print('loading val_valid_data_moments...')
self.val_index = DataIndex(lo(val_all_steer_path), -1, 0)
@staticmethod
def get_data(run_code, seg_num, offset):
data = Segment_Data.get_data(run_code, seg_num, offset,
ARGS.stride * ARGS.nsteps, offset,
ARGS.nframes, ignore=ARGS.ignore,
require_one=ARGS.require_one,
use_states=ARGS.use_states)
return data
@staticmethod
def next(data_index):
if data_index.ctr >= len(data_index.valid_data_moments) - (
1 + ARGS.batch_size): # Skip last batch if it runs out of data
data_index.ctr = -1
data_index.epoch_counter += 1
data_index.epoch_complete = True
if data_index.ctr == -1:
data_index.ctr = 0
print('shuffle start')
random.shuffle(data_index.valid_data_moments)
print('shuffle finished')
data_index.ctr += 1
return data_index.valid_data_moments[data_index.ctr]