1+ import json
2+
13import pytest
24import torch
35from pydantic .tools import parse_obj_as , schema_json_of
46
7+ from docarray import BaseDoc
58from docarray .base_doc .io .json import orjson_dumps
9+ from docarray .proto import DocProto
610from docarray .typing import TorchEmbedding , TorchTensor
711
812
13+ class MyDoc (BaseDoc ):
14+ tens : TorchTensor
15+
16+
917@pytest .mark .proto
1018def test_proto_tensor ():
1119 tensor = parse_obj_as (TorchTensor , torch .zeros (3 , 224 , 224 ))
@@ -63,7 +71,7 @@ def test_wrong_but_reshapable():
6371 parse_obj_as (TorchTensor [3 , 224 , 224 ], torch .zeros (224 , 224 ))
6472
6573
66- def test_inependent_variable_dim ():
74+ def test_independent_variable_dim ():
6775 # test independent variable dimensions
6876 tensor = parse_obj_as (TorchTensor [3 , 'x' , 'y' ], torch .zeros (3 , 224 , 224 ))
6977 assert isinstance (tensor , TorchTensor )
@@ -177,3 +185,53 @@ class MMdoc(BaseDoc):
177185
178186 doc_copy .embedding = torch .randn (32 )
179187 assert not (doc .embedding == doc_copy .embedding ).all ()
188+
189+
190+ @pytest .mark .parametrize ('requires_grad' , [True , False ])
191+ def test_json_serialization (requires_grad ):
192+ orig_doc = MyDoc (tens = torch .rand (10 , requires_grad = requires_grad ))
193+ serialized_doc = orig_doc .to_json ()
194+ assert serialized_doc
195+ assert isinstance (serialized_doc , str )
196+
197+ json_doc = json .loads (serialized_doc )
198+ assert json_doc ['tens' ]
199+ assert len (json_doc ['tens' ]) == 10
200+
201+
202+ @pytest .mark .parametrize ('protocol' , ['pickle' , 'protobuf' ])
203+ @pytest .mark .parametrize ('requires_grad' , [True , False ])
204+ def test_bytes_serialization (requires_grad , protocol ):
205+ orig_doc = MyDoc (tens = torch .rand (10 , requires_grad = requires_grad ))
206+ serialized_doc = orig_doc .to_bytes (protocol = protocol )
207+ assert serialized_doc
208+ assert isinstance (serialized_doc , bytes )
209+
210+ conv_doc = MyDoc .from_bytes (serialized_doc , protocol = protocol )
211+ assert isinstance (conv_doc .tens , TorchTensor )
212+ assert conv_doc .tens .shape == (10 ,)
213+
214+
215+ @pytest .mark .parametrize ('protocol' , ['pickle' , 'protobuf' ])
216+ @pytest .mark .parametrize ('requires_grad' , [True , False ])
217+ def test_base64_serialization (requires_grad , protocol ):
218+ orig_doc = MyDoc (tens = torch .rand (10 , requires_grad = requires_grad ))
219+ serialized_doc = orig_doc .to_base64 (protocol = protocol )
220+ assert serialized_doc
221+ assert isinstance (serialized_doc , str )
222+
223+ conv_doc = MyDoc .from_base64 (serialized_doc , protocol = protocol )
224+ assert isinstance (conv_doc .tens , TorchTensor )
225+ assert conv_doc .tens .shape == (10 ,)
226+
227+
228+ @pytest .mark .parametrize ('requires_grad' , [True , False ])
229+ def test_protobuf_serialization (requires_grad ):
230+ orig_doc = MyDoc (tens = torch .rand (10 , requires_grad = requires_grad ))
231+ serialized_doc = orig_doc .to_protobuf ()
232+ assert serialized_doc
233+ assert isinstance (serialized_doc , DocProto )
234+
235+ conv_doc = MyDoc .from_protobuf (serialized_doc )
236+ assert isinstance (conv_doc .tens , TorchTensor )
237+ assert conv_doc .tens .shape == (10 ,)
0 commit comments