MCPcopy Create free account
hub / github.com/AI-Hypercomputer/maxtext / load_next_batch

Method load_next_batch

src/MaxText/data_loader.py:106–157  ·  view source on GitHub ↗

Updates the batch size based on the schedule and then loads the next batch using the parent method.

(self, *args, rampup_manager=None, **kwargs)

Source from the content-addressed store, hash-verified

104 self.buffer_start = 0
105
106 def load_next_batch(self, *args, rampup_manager=None, **kwargs):
107 """
108 Updates the batch size based on the schedule and then loads the next
109 batch using the parent method.
110 """
111 # If ramp-up is not active, just behave like the parent
112 if not self.rampup_active:
113 return super().load_next_batch()
114
115 slice_start, slice_end = self.buffer_start, self.buffer_start + rampup_manager.global_batch_size_current
116
117 # Load new batch if batch_buffer is None
118 if self.batch_buffer is None:
119 self.batch_buffer = super().load_next_batch_pre_sharding()
120 slice_start, slice_end = 0, rampup_manager.global_batch_size_current
121
122 # If the slice end overpast batch end we collect new batch data
123 if slice_end > rampup_manager.global_batch_size_end:
124 old_buffer, self.batch_buffer = self.batch_buffer, super().load_next_batch_pre_sharding()
125
126 # self.global_batch_size_end is batch_buffer size
127 def _slice_and_concat(old_data, new_data):
128 sliced_old_data = jax.lax.dynamic_slice_in_dim(
129 old_data,
130 slice_start,
131 rampup_manager.global_batch_size_end - slice_start,
132 axis=0,
133 )
134 sliced_new_data = jax.lax.dynamic_slice_in_dim(
135 new_data,
136 0,
137 slice_end - rampup_manager.global_batch_size_end,
138 axis=0,
139 )
140 return jax.lax.concatenate((sliced_old_data, sliced_new_data), dimension=0)
141
142 self.buffer_start = slice_end - rampup_manager.global_batch_size_end
143 output = jax.tree.map(_slice_and_concat, old_buffer, self.batch_buffer)
144 else:
145
146 def _slice(data):
147 return jax.lax.dynamic_slice_in_dim(
148 data,
149 slice_start,
150 rampup_manager.global_batch_size_current,
151 axis=0,
152 )
153
154 self.buffer_start = slice_end
155 output = jax.tree.map(_slice, self.batch_buffer)
156 self.rampup_active = rampup_manager.update()
157 return maybe_shard_with_name(output, self.input_data_shardings, self.config.shard_mode)
158
159
160def create_dataloader(config, mesh, data_iterator, goodput_recorder, rampup_manager):

Calls 5

maybe_shard_with_nameFunction · 0.90
updateMethod · 0.80
load_next_batchMethod · 0.45
mapMethod · 0.45