@@ -30,11 +30,11 @@ class FlattenResult:
3030 dataframe : pd .DataFrame
3131 """The flattened DataFrame."""
3232
33- row_groups : dict [ str , list [int ]]
34- """
35- A mapping from original row index to the new row indices that were created
36- from it.
37- """
33+ row_labels : list [str ] | None
34+ """A list of original row labels for each row in the flattened DataFrame."""
35+
36+ continuation_rows : set [ int ] | None
37+ """A set of row indices that are continuation rows."""
3838
3939 cleared_on_continuation : list [str ]
4040 """A list of column names that should be cleared on continuation rows."""
@@ -50,7 +50,8 @@ def flatten_nested_data(
5050 if dataframe .empty :
5151 return FlattenResult (
5252 dataframe = dataframe .copy (),
53- row_groups = {},
53+ row_labels = None ,
54+ continuation_rows = None ,
5455 cleared_on_continuation = [],
5556 nested_columns = set (),
5657 )
@@ -77,15 +78,19 @@ def flatten_nested_data(
7778 if not array_columns :
7879 return FlattenResult (
7980 dataframe = result_df ,
80- row_groups = {},
81+ row_labels = None ,
82+ continuation_rows = None ,
8183 cleared_on_continuation = clear_on_continuation_cols ,
8284 nested_columns = nested_originated_columns ,
8385 )
8486
85- result_df , array_row_groups = _explode_array_columns (result_df , array_columns )
87+ result_df , row_labels , continuation_rows = _explode_array_columns (
88+ result_df , array_columns
89+ )
8690 return FlattenResult (
8791 dataframe = result_df ,
88- row_groups = array_row_groups ,
92+ row_labels = row_labels ,
93+ continuation_rows = continuation_rows ,
8994 cleared_on_continuation = clear_on_continuation_cols ,
9095 nested_columns = nested_originated_columns ,
9196 )
@@ -192,10 +197,10 @@ def _flatten_array_of_struct_columns(
192197
193198def _explode_array_columns (
194199 dataframe : pd .DataFrame , array_columns : list [str ]
195- ) -> tuple [pd .DataFrame , dict [str , list [int ] ]]:
200+ ) -> tuple [pd .DataFrame , list [str ], set [int ]]:
196201 """Explode array columns into new rows."""
197202 if not array_columns :
198- return dataframe , {}
203+ return dataframe , [], set ()
199204
200205 original_cols = dataframe .columns .tolist ()
201206 work_df = dataframe
@@ -243,7 +248,7 @@ def _explode_array_columns(
243248
244249 if not exploded_dfs :
245250 # This should not be reached if array_columns is not empty
246- return dataframe , {}
251+ return dataframe , [], set ()
247252
248253 # Merge the exploded columns
249254 merged_df = exploded_dfs [0 ]
@@ -260,22 +265,20 @@ def _explode_array_columns(
260265 drop = True
261266 )
262267
263- # Create row groups
264- array_row_groups = {}
268+ # Generate row labels and continuation mask efficiently
265269 grouping_col_name = (
266270 "_original_index" if original_index_name is None else original_index_name
267271 )
268- if grouping_col_name in merged_df .columns :
269- for orig_idx , group in merged_df .groupby (grouping_col_name ):
270- array_row_groups [str (orig_idx )] = group .index .tolist ()
272+ row_labels = merged_df [grouping_col_name ].astype (str ).tolist ()
273+ continuation_rows = set (merged_df .index [merged_df ["_row_num" ] > 0 ])
271274
272275 # Restore original columns
273276 result_df = merged_df [original_cols ]
274277
275278 if original_index_name :
276279 result_df = result_df .set_index (original_index_name )
277280
278- return result_df , array_row_groups
281+ return result_df , row_labels , continuation_rows
279282
280283
281284def _flatten_struct_columns (
0 commit comments