@@ -183,20 +183,23 @@ def append_transform(self, transform: BaseOperation):
183183 """
184184
185185
186- def _dict_input_fn (columns : Sequence [str ],
187- batch : Sequence [Dict [str , Any ]]) -> List [str ]:
186+ def _dict_input_fn (
187+ columns : Sequence [str ], batch : Sequence [Union [Dict [str , Any ],
188+ beam .Row ]]) -> List [str ]:
188189 """Extract text from specified columns in batch."""
190+ if batch and hasattr (batch [0 ], '_asdict' ):
191+ batch = [row ._asdict () if hasattr (row , '_asdict' ) else row for row in batch ]
192+
189193 if not batch or not isinstance (batch [0 ], dict ):
190194 raise TypeError (
191195 'Expected data to be dicts, got '
192196 f'{ type (batch [0 ])} instead.' )
193-
194197 result = []
195198 expected_keys = set (batch [0 ].keys ())
196199 expected_columns = set (columns )
197200 # Process one batch item at a time
198201 for item in batch :
199- item_keys = item .keys ()
202+ item_keys = item .keys () if isinstance ( item , dict ) else set ()
200203 if set (item_keys ) != expected_keys :
201204 extra_keys = item_keys - expected_keys
202205 missing_keys = expected_keys - item_keys
@@ -212,21 +215,31 @@ def _dict_input_fn(columns: Sequence[str],
212215
213216 # Get all columns for this item
214217 for col in columns :
215- result .append (item [col ])
218+ if isinstance (item , dict ):
219+ result .append (item [col ])
216220 return result
217221
218222
219223def _dict_output_fn (
220224 columns : Sequence [str ],
221- batch : Sequence [Dict [str , Any ]],
222- embeddings : Sequence [Any ]) -> List [ Dict [ str , Any ]]:
225+ batch : Sequence [Union [ Dict [str , Any ], beam . Row ]],
226+ embeddings : Sequence [Any ]) -> list [ Union [ dict [ str , Any ], beam . Row ]]:
223227 """Map embeddings back to columns in batch."""
228+ is_beam_row = False
229+ if batch and hasattr (batch [0 ], '_asdict' ):
230+ is_beam_row = True
231+ batch = [row ._asdict () if hasattr (row , '_asdict' ) else row for row in batch ]
232+
224233 result = []
225234 for batch_idx , item in enumerate (batch ):
226235 for col_idx , col in enumerate (columns ):
227236 embedding_idx = batch_idx * len (columns ) + col_idx
228- item [col ] = embeddings [embedding_idx ]
237+ if isinstance (item , dict ):
238+ item [col ] = embeddings [embedding_idx ]
229239 result .append (item )
240+
241+ if is_beam_row :
242+ result = [beam .Row (** item ) for item in result if isinstance (item , dict )]
230243 return result
231244
232245
0 commit comments