LIMA
Libre Multilingual Analyzer — C++ API
Loading...
Searching...
No Matches
torch_matrix.h
Go to the documentation of this file.
1// Copyright 2021 CEA LIST
2// SPDX-FileCopyrightText: 2022 CEA LIST <gael.de-chalendar@cea.fr>
3//
4// SPDX-License-Identifier: MIT
5
6#ifndef DEEPLIMA_SRC_INFERENCE_TORCH_VECTORIZER_H
7#define DEEPLIMA_SRC_INFERENCE_TORCH_VECTORIZER_H
8
9#include <torch/torch.h>
10
11namespace deeplima
12{
13
14template <class T>
16{
17public:
18 typedef T value_t;
19 typedef torch::Tensor tensor_t;
20
22 : m_start_time(0),
23 m_max_time(0),
24 m_max_feat(0),
26 m_accessor(m_tensor.accessor<T, 2>())
27 {}
28
29 TorchMatrix(int64_t max_time, int64_t max_feat)
30 : m_start_time(0),
31 m_max_time(max_time),
32 m_max_feat(max_feat),
33 m_tensor(create_tensor(max_time, max_feat)),
34 m_accessor(m_tensor.accessor<T, 2>())
35 {
36 assert(m_start_time >= 0);
37 assert(m_max_time > m_start_time);
38 //assert(m_max_feat > 0);
39 assert(m_max_feat < std::numeric_limits<int64_t>::max());
40 }
41
42 TorchMatrix(torch::Tensor input, int64_t start_time, int64_t max_time)
43 : m_start_time(start_time),
44 m_max_time(max_time),
45 m_max_feat(input.size(1)),
46 m_tensor(input),
47 m_accessor(m_tensor.accessor<T, 2>())
48 {
49 assert(m_start_time >= 0);
50 assert(m_max_time > m_start_time);
51 assert(m_max_feat > 0);
52 assert(m_max_feat < std::numeric_limits<int64_t>::max());
53 }
54
55 void init(int64_t max_time, int64_t max_feat)
56 {
57 assert(m_start_time >= 0);
58 assert(max_time > m_start_time);
59 assert(max_feat > 0);
60 assert(max_feat < std::numeric_limits<int64_t>::max());
61
62 m_max_time = max_time;
63 m_max_feat = max_feat;
64 m_tensor = create_tensor(max_time, max_feat);
65 m_accessor = m_tensor.accessor<T, 2>();
66 }
67
68 inline void set(int64_t time, uint64_t feat, T value)
69 {
70 assert(feat < std::numeric_limits<uint64_t>::max());
71 assert(time < std::numeric_limits<int64_t>::max() - m_start_time);
72 assert(time < m_tensor.size(0));
73
74 assert(0 == m_start_time);
75 m_accessor[m_start_time + time][feat] = value;
76 assert(0 == m_start_time);
77 }
78
79 inline uint64_t get(uint64_t time, uint64_t feat)
80 {
81 assert(feat < std::numeric_limits<uint64_t>::max());
82 assert(time < std::numeric_limits<uint64_t>::max() - m_start_time);
83
84 return m_accessor[m_start_time + time][feat];
85 }
86
87 inline uint64_t size() const
88 {
89 return m_max_time - m_start_time;
90 }
91
92 inline uint64_t get_max_feat() const
93 {
94 return m_max_feat;
95 }
96
97 const torch::Tensor& get_tensor() const
98 {
99 return m_tensor;
100 }
101
102 void to(torch::Device& device)
103 {
104 m_tensor.to(device);
105 }
106
107protected:
108
109 void create();
110 torch::Tensor create_empty_tensor()
111 {
112 return torch::zeros({0, 0}, torch::TensorOptions().dtype(torch::kInt64));
113 }
114 torch::Tensor create_tensor(int64_t max_time, int64_t max_feat);
115
117 int64_t m_max_time;
118 int64_t m_max_feat;
119
120 torch::Tensor m_tensor;
121 torch::TensorAccessor<T, 2> m_accessor;
122};
123
124template<>
125inline torch::Tensor TorchMatrix<int64_t>::create_tensor(int64_t max_time, int64_t max_feat)
126{
127 assert(max_time > 0);
128 //assert(max_feat > 0);
129
130 return torch::zeros({max_time, max_feat}, torch::TensorOptions().dtype(torch::kInt64));
131}
132
133template<>
134inline torch::Tensor TorchMatrix<float>::create_tensor(int64_t max_time, int64_t max_feat)
135{
136 assert(max_time > 0);
137 assert(max_feat > 0);
138
139 return torch::zeros({max_time, max_feat}, torch::TensorOptions().dtype(torch::kFloat32));
140}
141
142template <class D, class Matrix>
144{
145public:
146
148
149 inline void set(Matrix& target, uint64_t time, uint64_t feat, const typename D::value_t value)
150 {
151 assert(feat < m_dicts.size());
152 int k = m_dicts[feat]->get_idx(value);
153 target.set(time, feat, k);
154 }
155
156protected:
157 std::vector<D*> m_dicts;
158};
159
160template<>
162{
163 assert(m_max_time > 0);
164 assert(m_max_feat > 0);
165
166 m_tensor = torch::zeros({m_max_time, m_max_feat}, torch::TensorOptions().dtype(torch::kInt64));
167 m_accessor = m_tensor.accessor<int64_t, 2>();
168}
169
170template<>
172{
173 assert(m_max_time > 0);
174 assert(m_max_feat > 0);
175
176 m_tensor = torch::zeros({m_max_time, m_max_feat}, torch::TensorOptions().dtype(torch::kFloat32));
177 m_accessor = m_tensor.accessor<float, 2>();
178}
179
180} // namespace deeplima
181
182#endif
std::vector< D * > m_dicts
void set(Matrix &target, uint64_t time, uint64_t feat, const typename D::value_t value)
uint64_t get_max_feat() const
TorchMatrix(torch::Tensor input, int64_t start_time, int64_t max_time)
uint64_t get(uint64_t time, uint64_t feat)
torch::Tensor tensor_t
const torch::Tensor & get_tensor() const
torch::Tensor create_tensor(int64_t max_time, int64_t max_feat)
TorchMatrix(int64_t max_time, int64_t max_feat)
uint64_t size() const
void init(int64_t max_time, int64_t max_feat)
torch::Tensor create_empty_tensor()
void to(torch::Device &device)
torch::Tensor m_tensor
torch::TensorAccessor< T, 2 > m_accessor
void set(int64_t time, uint64_t feat, T value)