@@ -680,17 +680,27 @@ def __init__(
680680 self .binary = binary
681681
682682 def expand (self , pcoll ):
683- if 'file_naming' in self .kwargs :
683+ kwargs = dict (self .kwargs )
684+ if 'file_naming' in kwargs :
684685 dir , name = self .path , ''
685686 else :
686687 dir , name = io .filesystems .FileSystems .split (self .path )
688+ num_shards = kwargs .pop ('num_shards' , None )
689+ max_writers_per_bundle = kwargs .pop ('max_writers_per_bundle' , None )
690+ write_to_files_kwargs = {}
691+ if num_shards is not None :
692+ write_to_files_kwargs ['shards' ] = num_shards
693+ write_to_files_kwargs ['max_writers_per_bundle' ] = 0
694+ elif max_writers_per_bundle is not None :
695+ write_to_files_kwargs ['max_writers_per_bundle' ] = max_writers_per_bundle
696+
697+ file_naming = kwargs .pop ('file_naming' , fileio .default_file_naming (name ))
687698 return pcoll | fileio .WriteToFiles (
688699 path = dir ,
689- shards = self .kwargs .pop ('num_shards' , None ),
690- file_naming = self .kwargs .pop (
691- 'file_naming' , fileio .default_file_naming (name )),
700+ file_naming = file_naming ,
692701 sink = lambda _ : _WriteToPandasFileSink (
693- self .writer , self .args , self .kwargs , self .incremental , self .binary ))
702+ self .writer , self .args , kwargs , self .incremental , self .binary ),
703+ ** write_to_files_kwargs )
694704
695705
696706class _WriteToPandasFileSink (fileio .FileSink ):
0 commit comments