LIMA
Libre Multilingual Analyzer — C++ API
Loading...
Searching...
No Matches
iterable_dataset.h
Go to the documentation of this file.
1// Copyright 2022 CEA LIST
2// SPDX-FileCopyrightText: 2022 CEA LIST <gael.de-chalendar@cea.fr>
3//
4// SPDX-License-Identifier: MIT
5
6#ifndef DEEPLIMA_LIBS_TASKS_GRAPH_DP_ITERABLE_DATASET_H
7#define DEEPLIMA_LIBS_TASKS_GRAPH_DP_ITERABLE_DATASET_H
8
9#include <torch/torch.h>
10
11namespace deeplima
12{
13namespace graph_dp
14{
15namespace train
16{
17
19{
20public:
21 class Batch
22 {
23 public:
24 inline bool empty() const
25 {
26 return m_gold.size(m_batch_dim) == 0;
27 }
28
29 inline size_t get_batch_size() const
30 {
31 return m_gold.size(m_batch_dim);
32 }
33
34 Batch() { }
35
36 Batch(const torch::Tensor& trainable_input,
37 const torch::Tensor& frozen_input,
38 const torch::Tensor& gold,
39 size_t batch_dim=0)
40 : m_batch_dim(batch_dim),
44 {
45 }
46
47 inline const torch::Tensor& trainable_input() const
48 {
49 return m_trainable_input;
50 }
51
52 inline const torch::Tensor& frozen_input() const
53 {
54 return m_frozen_input;
55 }
56
57 inline const torch::Tensor& gold() const
58 {
59 return m_gold;
60 }
61
62 protected:
63 const size_t m_batch_dim = 0;
64
65 const torch::Tensor m_trainable_input;
66 const torch::Tensor m_frozen_input;
67 const torch::Tensor m_gold;
68 };
69 virtual ~BatchIterator() = default;
70 virtual void set_batch_size(int64_t batch_size) = 0;
71 virtual void start_epoch() = 0;
72 virtual bool end() = 0;
73 virtual const Batch next_batch() = 0;
74};
75
77{
78public:
79 virtual std::shared_ptr<BatchIterator> get_iterator() const = 0;
80};
81
82} // train
83} // graph_dp
84} // deeplima
85
86#endif // DEEPLIMA_LIBS_TASKS_GRAPH_DP_ITERABLE_DATASET_H
Batch(const torch::Tensor &trainable_input, const torch::Tensor &frozen_input, const torch::Tensor &gold, size_t batch_dim=0)
virtual void set_batch_size(int64_t batch_size)=0
virtual std::shared_ptr< BatchIterator > get_iterator() const =0