-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathmvp.js
More file actions
90 lines (83 loc) · 2.23 KB
/
Copy pathmvp.js
File metadata and controls
90 lines (83 loc) · 2.23 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
"use strict"
module.exports = matrixVectorProduct
function generateGetter(type, data, index) {
if(type === "generic") {
return [ data, ".get(", index, ")" ].join("")
} else {
return [ data, "[", index, "]" ].join("")
}
}
function generateSetter(type, data, index, value) {
if(type === "generic") {
return [ data, ".set(", index, ",", value, ")" ].join("")
} else {
return [ data, "[", index, "]=", value ].join("")
}
}
function generateProduct(rowMajor, typesig) {
var funcName = ["matrixVectorProd", rowMajor ? "RowMajor" : "ColumnMajor", typesig.join("t")].join("")
var code = ["function ", funcName, "(o,m,v){\
var o_d=o.data,o_i=o.offset,o_s=o.stride[0],\
m_d=m.data,m_i=m.offset,m_s0=m.stride[0],m_s1=m.stride[1],\
v_d=v.data,v_i=v.offset,v_s=v.stride[0],r=m.shape[0],c=m.shape[1],"]
if(rowMajor) {
code.push(
"s1=m_s0-m_s1*r;for(var j=0;j<c;++j){\
var v_p=v_i,x=0;\
for(var i=0;i<r;++i){\
x+=", generateGetter(typesig[1], "m_d", "m_i"), "*", generateGetter(typesig[2], "v_d", "v_p"), ";\
v_p+=v_s;\
m_i+=m_s1;\
}",
generateSetter(typesig[0], "o_d", "o_i", "x"), ";\
o_i+=o_s;\
m_i+=s1;\
}")
} else {
code.push("o_p=o_i;for(var i=0;i<r;++i){",
generateSetter(typesig[0], "o_d", "o_p", "0"), ";\
o_p+=o_s}\
var s0=m_s1-m_s0*c;for(var j=0;j<r;++j){\
o_p=o_i;",
"var x=", generateGetter(typesig[2], "v_d", "v_i"), ";\
for(var i=0;i<c;++i){")
if(typesig[0] === "generic") {
code.push("o_d.set(o_p,o_d.get(o_p)+x*", generateGetter(typesig[1], "m_d", "m_i"), ";")
} else {
code.push("o_d[o_p]+=", generateGetter(typesig[1], "m_d", "m_i"), "*x;")
}
code.push("m_i+=m_s0;\
o_p+=o_s;\
}\
m_i+=s0;\
v_i+=v_s;\
}")
}
code.push("}return ", funcName, ";")
//Compile result
var proc = new Function(code.join(""))
return proc()
}
var ROW_CACHE = {}
var COLUMN_CACHE = {}
function matrixVectorProduct(out, m, v) {
var cache
if(m.stride[0] > m.stride[1]) {
cache = ROW_CACHE
} else {
cache = COLUMN_CACHE
}
var a = cache[out.dtype]
if(!a) {
cache[out.dtype] = a = {}
}
var b = a[m.dtype]
if(!b) {
a[m.dtype] = b = {}
}
var c = b[v.dtype]
if(!c) {
b[v.dtype] = c = generateProduct(m.stride[0] > m.stride[1], [out.dtype, m.dtype, v.dtype])
}
c(out,m,v)
}