Uh oh!
There was an error while loading. Please reload this page.
- Notifications
You must be signed in to change notification settings - Fork 84
Expand file tree
/
Copy pathSimpleColorTransform.lua
More file actions
Latest commit
90 lines (79 loc) · 3.26 KB
/
Copy pathSimpleColorTransform.lua
File metadata and controls
90 lines (79 loc) · 3.26 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
--[[
Simple Color transformation module: This module implements a simple data
augmentation technique of changing the pixel values of input image by adding
sample sampled small quantities.
Works only
--]]
localSimpleColorTransform, Parent=torch.class('nn.SimpleColorTransform', 'nn.Module')
functionSimpleColorTransform:__init(inputChannels, range)
Parent.__init(self)
self.train=true
self.inputChannels=inputChannels
assert(inputChannels==range:nElement(),
"Number of input channels and number of range values don't match.")
self.range=range
end
functionSimpleColorTransform:updateOutput(input)
self.output:resizeAs(input):copy(input)
ifself.trainthen
self.noise=self.noiseorself.output.new()
self._tempNoise=self._tempNoiseorself.output.new()
self._tempNoiseExpanded=self._tempNoiseExpandedorself.output.new()
self._tempNoiseSamples=self._tempNoiseSamplesorself.output.new()
ifself.output:nDimension() ==4then
localbatchSize=self.output:size(1)
localchannels=self.output:size(2)
localheight=self.output:size(3)
localwidth=self.output:size(4)
assert(channels==self.inputChannels)
-- Randomly sample noise for each channel
self.noise:resize(batchSize, channels)
fori=1, channelsdo
self.noise[{{}, {i}}]:uniform(-self.range[i], self.range[i])
end
self._tempNoise=self.noise:view(batchSize, self.inputChannels, 1, 1)
self._tempNoiseExpanded:expand(self._tempNoise, batchSize,
channels, height, width)
self._tempNoiseSamples:resizeAs(self._tempNoiseExpanded)
:copy(self._tempNoiseExpanded)
self.output:add(self._tempNoiseSamples)
elseifself.output:nDimension() ==3then
localchannels=self.output:size(1)
localheight=self.output:size(2)
localwidth=self.output:size(3)
assert(channels==self.inputChannels)
-- Randomly sample noise for each channel
self.noise:resize(channels)
fori=1, channelsdo
self.noise[i] =torch.uniform(-self.range[i], self.range[i])
end
self._tempNoise=self.noise:view(self.inputChannels, 1, 1)
self._tempNoiseExpanded:expand(self._tempNoise, channels,
height, width)
self._tempNoiseSamples:resizeAs(self._tempNoiseExpanded)
:copy(self._tempNoiseExpanded)
self.output:add(self._tempNoiseSamples)
else
error("Invalid input dimensionality.")
end
end
returnself.output
end
functionSimpleColorTransform:updateGradInput(input, gradOutput)
ifself.trainthen
self.gradInput:resizeAs(gradOutput):copy(gradOutput)
else
error('backprop only defined while training')
end
returnself.gradInput
end
functionSimpleColorTransform:type(type, tensorCache)
self.noise=nil
self._tempNoise=nil
self._tempNoiseExpanded=nil
self._tempNoiseSamples=nil
Parent.type(self, type, tensorCache)
end
functionSimpleColorTransform:__tostring__()
returnstring.format('SimpleColorTransform', torch.type(self))
end