- Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathconvolution.java
More file actions
Latest commit
112 lines (101 loc) · 2.76 KB
/
Copy pathconvolution.java
File metadata and controls
112 lines (101 loc) · 2.76 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
packagetensordef;
importbasicops.*;
publicclassconvolutionextendssuperopdef
{
tensorarray3darr;
tensorarrayfilters[];
tensorgraphgraph;
tensorarray3dpaddedinputs;
tensorarray3deval;
intnumfilters;
intpadding;
intfiltersize;
tensorarrayback[][][];
tensorarrayopsops;
tensorarray3deval1[][][];
backpropagationstructure<convolution> curstruct;
dot3dmulops[][][];
dotdotops[][][];
reduce_sum3dredops[][][];
tensorarrayipslice[][];
publicconvolution(tensorarray3darr,intfiltersize,intnumfilters,tensorgraphgraph,Stringpad)
{
padding=0;
ops=newtensorarrayops();
this.arr=arr;
this.numfilters=numfilters;
this.graph=graph;
this.filtersize=filtersize;
if(pad.equals("SAME"))
{
padding=(filtersize-1)/2;
paddedinputs=ops.pad(arr,arr.dim1+2*padding,arr.dim2+2*padding);
}else
{
paddedinputs=arr;
}
//System.out.println(filtersize);
eval=newtensorarray3d((arr.dim1-filtersize+2*padding)+1,arr.dim2-filtersize+2*padding+1,numfilters,false);
back=newtensorarray[(arr.dim1-filtersize+2*padding)+1][arr.dim2-filtersize+2*padding+1][numfilters];
ipslice=newtensorarray[eval.dim1][eval.dim2];
filters=newtensorarray[numfilters];
//
for(inti=0;i<numfilters;i++)
{
filters[i]=newtensorarray(1,filtersize*filtersize*arr.dim3,true);
}
//filters[0].print();
dotops=newdot[eval.dim1][eval.dim2][numfilters];
for(inti=0;i<=paddedinputs.dim1-filtersize;i++)
{
for(intj=0;j<=paddedinputs.dim2-filtersize;j++)
{
ipslice[i][j]=ops.stretch(ops.getslices(paddedinputs,i,i+filtersize,j,j+filtersize),true);
for(intk=0;k<numfilters;k++)
{
//[k].print();
//System.out.println(k);
back[i][j][k]=newtensorarray(1,1,eval.trainable);
//System.out.println(back[i][j][k]);
dotops[i][j][k]=newdot(filters[k],ipslice[i][j],graph);
}
}
}
curstruct=newbackpropagationstructure<convolution>(this,null,eval);
graph.addtolist(curstruct);
}
publictensorarray3dforwardconv()
{
for(inti=0;i<=paddedinputs.dim1-filtersize;i++)
{
for(intj=0;j<=paddedinputs.dim2-filtersize;j++)
{
for(intk=0;k<numfilters;k++)
{
eval.arr[i][j][k].data=dotops[i][j][k].forward().arr[0][0].data;
}
}
}
// System.out.println("hello");
returneval;
}
publicvoidbackwardconv(tensorarray3dbackflow)
{
//System.out.println("bcfgsg");
//backflow.print();
ops.tensorarray3dtoarrayoftensorarray2d(backflow,back);
for(inti=0;i<=paddedinputs.dim1-filtersize;i++)
{
for(intj=0;j<=paddedinputs.dim2-filtersize;j++)
{
for(intk=0;k<numfilters;k++)
{
//System.out.println(back[i][j][k]);
//System.out.println("----------------");
dotops[i][j][k].backward(back[i][j][k]);
}
}
}
graph.removefromlist(curstruct);
}
}