A concrete implementation of the DataSerializer interface that serializes and deserializes data to/from the FlatTensor format.
| 280 | |
| 281 | |
| 282 | class 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( |
no outgoing calls