-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathnumlua.lua
More file actions
130 lines (94 loc) · 3.06 KB
/
Copy pathnumlua.lua
File metadata and controls
130 lines (94 loc) · 3.06 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
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
local Class = require('class')
local deepcopy = require('deepcopy')
--- @module
local numlua = {}
numlua.Array = Class.new()
function numlua.Array:init()
self._data = {}
self._strides = {}
self._startOffsets = setmetatable({}, {__index = function() return 0 end}) -- positive number
self._stopOffsets = setmetatable({}, {__index = function() return 0 end}) -- positive number
-- TODO:
-- - generate self.shape somewhere
-- - set ndim
end
function numlua.Array:_getShape()
-- compute shape from strides and offsets
local outShape = {}
local remainingLen = #self._data
for i, stride in ipairs(self._strides) do
local startOffset = self._startOffsets[i] * stride
local stopOffset = self._stopOffsets[i] * stride
local lenAxis = (remainingLen / stride) - startOffset - stopOffset
table.insert(outShape, lenAxis)
remainingLen = remainingLen / lenAxis
end
return outShape
end
function numlua.Array:ndim()
return #self._strides
end
function numlua.Array:view()
-- return new view on data
local outArr = numlua.Array()
-- view on data
outArr._data = self._data
-- copy metainformation
outArr._strides = deepcopy(self._strides)
outArr._startOffsets = deepcopy(self._startOffsets)
outArr._stopOffsets = deepcopy(self._stopOffsets)
outArr._owndata = false
return outArr
end
function numlua.Array:_getByPositiveIndex(...)
--[[
Get single data point by ṕositive index.
TODO: remove checks and put them in higher API method (performance)
]]
local indices = {...}
if #indices ~= self:ndim() then
error("Number of indices must match number of dimensions")
end
local shape = self:_getShape()
for i, index in ipairs(indices) do
if index < 1 then
error("Provide indices > 0")
end
if index > shape[i] then
error("Index " .. tostring(index) .. " too big for axis " .. tostring(i))
end
end
local totIndex = 0
for i, index in ipairs(indices) do
-- calculate in 0-indexing
totIndex = totIndex + self._strides[i] * (self._startOffsets[i] + index - 1)
end
-- convert to 1-indexing (kill me pls)
totIndex = totIndex + 1
return self._data[totIndex]
end
function numlua.Array:reshape(...)
local newShape = {...}
end
function numlua.Array:T()
if self:ndim() ~= 2 then
error(":T() only works with 2D arrays")
end
local outArr = self:view()
outArr._strides[1], outArr._strides[2] = outArr._strides[2], outArr._strides[1]
outArr._startOffsets[1], outArr._startOffsets[2] = outArr._startOffsets[2], outArr._startOffsets[1]
outArr._stopOffsets[1], outArr._stopOffsets[2] = outArr._stopOffsets[2], outArr._stopOffsets[1]
return outArr
end
-- functions to generate arrays
function numlua.toArray(t)
end
function numlua.zeros(...)
end
function numlua.ones(...)
end
function numlua.range(start, stop, step)
end
function numlua.linspace(start, stop, numSteps, endpoint)
end
return numlua