-
Notifications
You must be signed in to change notification settings - Fork 7
Expand file tree
/
Copy pathMapReduce.lua
More file actions
84 lines (62 loc) · 2.75 KB
/
Copy pathMapReduce.lua
File metadata and controls
84 lines (62 loc) · 2.75 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
local MapReduce, parent = torch.class('nn.MapReduce', 'nn.Container')
function MapReduce:__init(mapper,reducer)
parent.__init(self)
self.mapper = mapper
self.reducer = reducer
self.modules = {}
table.insert(self.modules,mapper)
table.insert(self.modules,reducer)
end
function MapReduce:updateOutput(input)
--first, reshape the data by pulling the second dimension into the first
self.inputSize = input:size()
local numPerExample = self.inputSize[2]
local minibatchSize = self.inputSize[1]
self.sizes = self.sizes or torch.LongStorage(self.inputSize:size() -1)
self.sizes[1] = minibatchSize*numPerExample
for i = 2,self.sizes:size() do
self.sizes[i] = self.inputSize[i+1]
end
self.reshapedInput = input:view(self.sizes)
self.mapped = self.mapper:updateOutput(self.reshapedInput)
self.sizes3 = self.mapped:size()
self.sizes2 = self.sizes2 or torch.LongStorage(self.mapped:dim() + 1)
self.sizes2[1] = minibatchSize
self.sizes2[2] = numPerExample
for i = 2,self.mapped:dim() do
self.sizes2[i+1] = self.mapped:size(i)
end
self.mappedAndReshaped = self.mapped:view(self.sizes2)
self.output = self.reducer:updateOutput(self.mappedAndReshaped)
return self.output
end
function MapReduce:backward(input,gradOutput)
local function operator(module,input,gradOutput) return module:backward(input,gradOutput) end
return self:genericBackward(operator,input,gradOutput)
end
function MapReduce:updateGradInput(input,gradOutput)
local function operator(module,input,gradOutput) return module:updateGradInput(input,gradOutput) end
return self:genericBackward(operator,input,gradOutput)
end
function MapReduce:accUpdateGradParameters(input,gradOutput,lr)
local function operator(module,input,gradOutput) return module:accUpdateGradParameters(input,gradOutput,lr) end
return self:genericBackward(operator,input,gradOutput)
end
function MapReduce:accGradParameters(input,gradOutput,lr)
local function operator(module,input,gradOutput) return module:accGradParameters(input,gradOutput,lr) end
return self:genericBackward(operator,input,gradOutput)
end
function MapReduce:genericBackward(operator,input, gradOutput)
operator(self.reducer,self.mappedAndReshaped,gradOutput)
local reducerGrad = self.reducer.gradInput
if reducerGrad:isContiguous() then
self.reshapedReducerGrad = reducerGrad:view(self.sizes3)
else
self.reshapedReducerGrad = self.reshapedReducerGrad or reducerGrad:clone()
self.reshapedReducerGrad:resizeAs(reducerGrad):copy(reducerGrad):resize(self.sizes3)
end
operator(self.mapper,self.reshapedInput,self.reshapedReducerGrad)
local mapperGrad = self.mapper.gradInput
self.gradInput = (mapperGrad:dim() > 0) and mapperGrad:view(self.inputSize) or nil --some modules return nil from backwards, such as the lookup table
return self.gradInput
end