MCPcopy Create free account
hub / github.com/pytorch/executorch / FlatTensorSerializer

Class FlatTensorSerializer

extension/flat_tensor/serialize/serialize.py:282–446  ·  view source on GitHub ↗

A concrete implementation of the DataSerializer interface that serializes and deserializes data to/from the FlatTensor format.

Source from the content-addressed store, hash-verified

280
281
282class FlatTensorSerializer(DataSerializer):
283 """A concrete implementation of the DataSerializer interface that
284 serializes and deserializes data to/from the FlatTensor format.
285 """
286
287 def __init__(self, config: Optional[FlatTensorConfig] = None) -> None:
288 """FlatTensorConfig holds information required for serialization,
289 eg. alignment.
290 """
291 if config is None:
292 self.config: FlatTensorConfig = FlatTensorConfig()
293 else:
294 self.config: FlatTensorConfig = config
295
296 def serialize(
297 self,
298 data: DataPayload,
299 ) -> Cord:
300 """Serializes a list of tensors and named data into a blob."""
301
302 segments: List[AlignedData] = []
303
304 # Add a config to place tensors in a single segment.
305 named_data = _extract_named_data(data, segments)
306
307 data_segments: List[DataSegment] = []
308 aggregated_segment_data = Cord()
309 for segment in segments:
310 prev_end = (
311 (data_segments[-1].offset + data_segments[-1].size)
312 if data_segments
313 else 0
314 )
315 alignment = math.lcm(self.config.segment_alignment, segment.alignment)
316 data_segments.append(
317 DataSegment(
318 offset=aligned_size(prev_end, alignment),
319 size=len(segment.data),
320 )
321 )
322 # Pad aggregated_segment_data to segment alignment.
323 segment_pad_length = padding_required(
324 len(aggregated_segment_data), alignment
325 )
326 if segment_pad_length > 0:
327 aggregated_segment_data.append(b"\x00" * segment_pad_length)
328 aggregated_segment_data.append(segment.data)
329
330 # Create FlatTensor, which describes of the contents of the file and
331 # points to all the data segments. It will be serialized to flatbuffer.
332 flat_tensor = FlatTensor(
333 version=_FLAT_TENSOR_VERSION,
334 segments=data_segments,
335 named_data=named_data,
336 )
337
338 flatbuffer_payload = _serialize_to_flatbuffer(flat_tensor)
339 padded_header_length: int = aligned_size(

Callers 5

__init__Method · 0.90
__init__Method · 0.90
test_round_tripMethod · 0.90

Calls

no outgoing calls