FazBrowse GitHub Viewer
|
Trending
|
URL:
|
Home
Tools:
[Download Repo ZIP]
[View Raw Code]
[Original HTTPS Page]
docarray/docarray/data/torch_dataset.py at main · docarray/docarray · GitHub
Uh oh!
There was an error while loading.
Please reload this page
.
docarray
/
docarray
Public
Notifications
You must be signed in to change notification settings
Fork
243
Star
3.1k
Code
Issues
68
Pull requests
43
Discussions
Actions
Security and quality
0
Insights
Additional navigation options
Code
Issues
Pull requests
Discussions
Actions
Security and quality
Insights
Expand file tree
Breadcrumbs
docarray
/
docarray
/
data
/
torch_dataset.py
Copy path
More file actions
More file actions
Latest commit
History
History
History
161 lines (118 loc) · 4.73 KB
Breadcrumbs
docarray
/
docarray
/
data
/
torch_dataset.py
Copy path
File metadata and controls
161 lines (118 loc) · 4.73 KB
Raw
Copy raw file
Download raw file
Open symbols panel
Edit and raw actions
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
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
from
typing
import
Callable
,
Dict
,
Generic
,
List
,
Optional
,
Type
,
TypeVar
from
torch
.
utils
.
data
import
Dataset
from
docarray
import
BaseDoc
,
DocList
,
DocVec
from
docarray
.
typing
import
TorchTensor
from
docarray
.
utils
.
_internal
.
_typing
import
change_cls_name
,
safe_issubclass
T_doc
=
TypeVar
(
'T_doc'
,
bound
=
BaseDoc
)
class
MultiModalDataset
(
Dataset
,
Generic
[
T_doc
]):
"""
A dataset that can be used inside a PyTorch DataLoader.
In other words, it implements the PyTorch Dataset interface.
The preprocessing dictionary passed to the constructor consists of keys that are
field names and values that are functions that take a single argument and return
a single argument.
---
```python
from torch.utils.data import DataLoader
from docarray import DocList
from docarray.data import MultiModalDataset
from docarray.documents import TextDoc
def prepend_number(text: str):
return f"Number {text}"
docs = DocList[TextDoc](TextDoc(text=str(i)) for i in range(16))
ds = MultiModalDataset[TextDoc](docs, preprocessing={'text': prepend_number})
loader = DataLoader(ds, batch_size=4, collate_fn=MultiModalDataset[TextDoc].collate_fn)
for batch in loader:
print(batch.text)
```
---
Nested fields can be accessed by using dot notation.
The document itself can be accessed using the empty string as the key.
Transformations that operate on reference types (such as Documents) can optionally
not return a value.
The transformations will be applied according to their order in the dictionary.
---
```python
import torch
from torch.utils.data import DataLoader
from docarray import DocList, BaseDoc
from docarray.data import MultiModalDataset
from docarray.documents import TextDoc
class Thesis(BaseDoc):
title: TextDoc
class Student(BaseDoc):
thesis: Thesis
def embed_title(title: TextDoc):
title.embedding = torch.ones(4)
def normalize_embedding(thesis: Thesis):
thesis.title.embedding = thesis.title.embedding / thesis.title.embedding.norm()
def add_nonsense(student: Student):
student.thesis.title.embedding = student.thesis.title.embedding + int(
student.thesis.title.text
)
docs = DocList[Student](Student(thesis=Thesis(title=str(i))) for i in range(16))
ds = MultiModalDataset[Student](
docs,
preprocessing={
"thesis.title": embed_title,
"thesis": normalize_embedding,
"": add_nonsense,
},
)
loader = DataLoader(ds, batch_size=4, collate_fn=ds.collate_fn)
for batch in loader:
print(batch.thesis.title.embedding)
```
---
:param docs: the `DocList` to be used as the dataset
:param preprocessing: a dictionary of field names and preprocessing functions
"""
doc_type
:
Optional
[
Type
[
BaseDoc
]]
=
None
__typed_ds__
:
Dict
[
Type
[
BaseDoc
],
Type
[
'MultiModalDataset'
]]
=
{}
def
__init__
(
self
,
docs
:
'DocList[T_doc]'
,
preprocessing
:
Dict
[
str
,
Callable
]
)
->
None
:
self
.
docs
=
docs
self
.
_preprocessing
=
preprocessing
def
__len__
(
self
):
return
len
(
self
.
docs
)
def
__getitem__
(
self
,
item
:
int
):
doc
=
self
.
docs
[
item
].
copy
(
deep
=
True
)
for
field
,
preprocess
in
self
.
_preprocessing
.
items
():
if
len
(
field
)
==
0
:
doc
=
preprocess
(
doc
)
or
doc
else
:
acc_path
=
field
.
split
(
'.'
)
_field_ref
=
doc
for
attr
in
acc_path
[:
-
1
]:
_field_ref
=
getattr
(
_field_ref
,
attr
)
attr
=
acc_path
[
-
1
]
value
=
getattr
(
_field_ref
,
attr
)
setattr
(
_field_ref
,
attr
,
preprocess
(
value
)
or
value
)
return
doc
@
classmethod
def
collate_fn
(
cls
,
batch
:
List
[
T_doc
]):
doc_type
=
cls
.
doc_type
if
doc_type
:
batch_da
=
DocVec
[
doc_type
](
# type: ignore
batch
,
tensor_type
=
TorchTensor
,
)
else
:
batch_da
=
DocVec
(
batch
,
tensor_type
=
TorchTensor
)
return
batch_da
@
classmethod
def
__class_getitem__
(
cls
,
item
:
Type
[
BaseDoc
])
->
Type
[
'MultiModalDataset'
]:
if
not
safe_issubclass
(
item
,
BaseDoc
):
raise
ValueError
(
f'
{
cls
.
__name__
}
[item] item should be a Document not a
{
item
}
'
)
if
item
not
in
cls
.
__typed_ds__
:
global
_TypedDataset
class
_TypedDataset
(
cls
):
# type: ignore
doc_type
=
item
change_cls_name
(
_TypedDataset
,
f'
{
cls
.
__name__
}
[
{
item
.
__name__
}
]'
,
globals
()
)
cls
.
__typed_ds__
[
item
]
=
_TypedDataset
return
cls
.
__typed_ds__
[
item
]
Back
|
FazBrowse Home
|
New Git URL