forked from LuZhenHuan/ECG-Classification-Demo
-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathecgDataProcess.lua
More file actions
61 lines (46 loc) · 1.51 KB
/
Copy pathecgDataProcess.lua
File metadata and controls
61 lines (46 loc) · 1.51 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
require 'torch'
local utils = require 'util.utils'
local DataLoader = torch.class('DataLoader')
function DataLoader:__init(kwargs)
local trainset = torch.load('RnnTrain1Dcut.t7')
local valset = torch.load('RnnVal1Dcut.t7')
local testset = torch.load('RnnTest1D.t7')
self.batch_size = utils.get_kwarg(kwargs, 'batch_size')
self.seq_length = utils.get_kwarg(kwargs, 'seq_length')
local N, T = self.batch_size, self.seq_length
-- Just slurp all the data into memory
local splits = {}
splits.train = trainset
splits.val = valset
splits.test = testset
self.x_splits = {}
self.y_splits = {}
self.split_sizes = {}
for split, v in pairs(splits) do
local num = v:nElement()
local extra = num % (N * T)
-- Ensure that `vy` is non-empty
if extra == 0 then
extra = N * T
end
-- Chop out the extra bits at the end to make it evenly divide
local vx = v[{{1, num - extra}}]:view(N, -1, T):transpose(1, 2):clone()
local vy = v[{{2, num - extra + 1}}]:view(N, -1, T):transpose(1, 2):clone()
self.x_splits[split] = vx
self.y_splits[split] = vy
self.split_sizes[split] = vx:size(1)
end
self.split_idxs = {train=1, val=1, test=1}
end
function DataLoader:nextBatch(split)
local idx = self.split_idxs[split]
assert(idx, 'invalid split ' .. split)
local x = self.x_splits[split][idx]
local y = self.y_splits[split][idx]
if idx == self.split_sizes[split] then
self.split_idxs[split] = 1
else
self.split_idxs[split] = idx + 1
end
return x, y
end