Skip to content

Commit f333c7d

Browse files
committed
Merge branch 'main' into jax_0.8.0_py3.12_v7x
2 parents db63dc4 + dcc3bcd commit f333c7d

16 files changed

Lines changed: 334 additions & 253 deletions

File tree

Dockerfile

Lines changed: 4 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -25,8 +25,6 @@ RUN echo "deb [signed-by=/usr/share/keyrings/cloud.google.gpg] https://packages.
2525
RUN mkdir -p /root
2626
WORKDIR /root
2727
# Introduce the minimum set of files for install.
28-
COPY README.md README.md
29-
COPY pyproject.toml pyproject.toml
3028
RUN mkdir axlearn && touch axlearn/__init__.py
3129
# Setup venv to suppress pip warnings.
3230
ENV VIRTUAL_ENV=/opt/venv
@@ -44,6 +42,7 @@ RUN pip install -qq --upgrade pip && \
4442
# Leverage multi-stage build for unit tests.
4543
FROM base AS ci
4644

45+
COPY pyproject.toml README.md /root/
4746
# TODO(markblee): Remove gcp,vertexai_tensorboard from CI.
4847
RUN uv pip install -qq .[core,audio,orbax,dev,gcp,vertexai_tensorboard] && \
4948
uv cache clean
@@ -76,6 +75,7 @@ FROM base AS dataflow
7675
# Beam workers default to creating a new virtual environment on startup. Instead, we want them to
7776
# pickup the venv setup above. An alternative is to install into the global environment.
7877
ENV RUN_PYTHON_SDK_IN_DEFAULT_ENVIRONMENT=1
78+
COPY pyproject.toml README.md /root/
7979
RUN uv pip install -qq .[core,gcp,dataflow] && uv cache clean
8080
COPY . .
8181

@@ -98,6 +98,7 @@ ARG INSTALL_PATHWAYS_JAXLIB=false
9898

9999
# Ensure we install the TPU version, even if building locally.
100100
# Jax will fallback to CPU when run on a machine without TPU.
101+
COPY pyproject.toml README.md /root/
101102
RUN uv pip install -qq --prerelease=allow .[core,tpu] && uv cache clean
102103
RUN if [ -n "$EXTRAS" ]; then uv pip install -qq .[$EXTRAS] && uv cache clean; fi
103104
RUN if [ "$INSTALL_PATHWAYS_JAXLIB" = "true" ]; then \
@@ -119,6 +120,7 @@ RUN curl -o cuda-keyring_1.1-1_all.deb https://developer.download.nvidia.com/com
119120
dpkg -i cuda-keyring_1.1-1_all.deb && \
120121
apt-get update && apt-get install -y cuda-libraries-dev-12-9 ibverbs-utils && \
121122
apt clean -y
123+
COPY pyproject.toml README.md /root/
122124
RUN uv pip install --prerelease=allow .[core,gpu] && uv cache clean
123125
COPY . .
124126

axlearn/cloud/gcp/lws_utils.py

Lines changed: 1 addition & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -7,7 +7,7 @@
77
from absl import flags
88

99
from axlearn.cloud.common.bundler import Bundler
10-
from axlearn.cloud.common.utils import AcceleratorConfig, FlagConfigurable, accelerator_flags
10+
from axlearn.cloud.common.utils import FlagConfigurable
1111
from axlearn.cloud.gcp.config import gcp_settings
1212
from axlearn.cloud.gcp.jobset_utils import TPUJobBuilder
1313
from axlearn.cloud.gcp.system_characteristics import USER_FACING_NAME_TO_SYSTEM_CHARACTERISTICS
@@ -27,7 +27,6 @@ class Config(FlagConfigurable.Config):
2727
Attributes:
2828
name: Name of the LeaderWorkerSet
2929
command: Command to be executed.
30-
accelerator: Accelerator configuration.
3130
env_vars: Optional env vars to set.
3231
service_account: Optional service account to execute the job as.
3332
output_dir: An optional GCS path to upload LWS outputs to.
@@ -37,7 +36,6 @@ class Config(FlagConfigurable.Config):
3736
# TODO: Change this to be a list of str[], to support different commands
3837
# between leader and workers
3938
command: Required[str] = REQUIRED
40-
accelerator: AcceleratorConfig = AcceleratorConfig()
4139
env_vars: dict[str, str] = {}
4240
service_account: Optional[str] = None
4341
output_dir: Optional[str] = None
@@ -47,7 +45,6 @@ class Config(FlagConfigurable.Config):
4745
def define_flags(cls, fv):
4846
super().define_flags(fv)
4947
common_kwargs = dict(flag_values=fv, allow_override=True)
50-
accelerator_flags(**common_kwargs)
5148
# NOTE: the parent typically sets these flags, so we leave them as None.
5249
flags.DEFINE_string("name", None, "Name of the LWS.", **common_kwargs)
5350
flags.DEFINE_string("command", None, "Command to execute.", **common_kwargs)
@@ -77,7 +74,6 @@ def from_flags(cls, fv: flags.FlagValues, **kwargs):
7774
cfg.service_account = cfg.service_account or gcp_settings(
7875
"k8s_service_account", default="default", fv=fv
7976
)
80-
cfg.accelerator.set(instance_type=fv.instance_type, num_replicas=fv.num_replicas)
8177
return cfg
8278

8379
def __init__(self, cfg: Config, *, bundler: Bundler):

axlearn/cloud/gcp/measurement.py

Lines changed: 18 additions & 18 deletions
Original file line numberDiff line numberDiff line change
@@ -30,16 +30,16 @@
3030
from ml_goodput_measurement import monitoring as goodput_monitoring
3131

3232
from axlearn.cloud.common.utils import parse_kv_flags, to_bool
33-
from axlearn.common import measurement
33+
from axlearn.common import measurement_base
3434
from axlearn.common.config import REQUIRED, Required, config_class, maybe_set_config
3535

3636

37-
@measurement.register_recorder("goodput")
38-
class GoodputRecorder(measurement.Recorder):
37+
@measurement_base.register_recorder("goodput")
38+
class GoodputRecorder(measurement_base.Recorder):
3939
"""Records overall training goodput."""
4040

4141
@config_class
42-
class Config(measurement.Recorder.Config):
42+
class Config(measurement_base.Recorder.Config):
4343
"""Configures GoodputRecorder.
4444
4545
Attributes:
@@ -77,7 +77,7 @@ def from_flags(cls, fv: flags.FlagValues) -> "GoodputRecorder":
7777
- jax_backend: The type of jax backend.
7878
- enable_monitoring: Boolean to enable/disable goodput monitoring (default: true).
7979
"""
80-
cfg: measurement.Recorder.Config = cls.default_config()
80+
cfg: measurement_base.Recorder.Config = cls.default_config()
8181
parsed_flags = parse_kv_flags(fv.recorder_spec, delimiter="=")
8282
if "upload_interval" in parsed_flags:
8383
parsed_flags["upload_interval"] = int(parsed_flags["upload_interval"])
@@ -100,7 +100,7 @@ def __init__(self, cfg):
100100
self._logger_name = f"goodput_logger_{cfg.name}"
101101

102102
@contextlib.contextmanager
103-
def record_event(self, event: measurement.EventType, *args, **kwargs):
103+
def record_event(self, event: measurement_base.EventType, *args, **kwargs):
104104
"""Records a goodput event using a context manager."""
105105
# Lazily instantiate the recorder if it hasn't been already.
106106
if self._recorder is None:
@@ -225,7 +225,7 @@ def monitor_goodput():
225225

226226
return monitor_goodput()
227227

228-
def record(self, event: measurement.Event, *args, **kwargs):
228+
def record(self, event: measurement_base.Event, *args, **kwargs):
229229
"""Deprecated: `record()` is not used in GoodputRecorder.
230230
Use the record_event context manager instead.
231231
"""
@@ -238,27 +238,27 @@ def record(self, event: measurement.Event, *args, **kwargs):
238238
logging_enabled=(jax.process_index() == 0),
239239
)
240240

241-
if event == measurement.Event.START_JOB:
241+
if event == measurement_base.Event.START_JOB:
242242
self._recorder.record_job_start_time(*args, **kwargs)
243-
elif event == measurement.Event.END_JOB:
243+
elif event == measurement_base.Event.END_JOB:
244244
self._recorder.record_job_end_time(*args, **kwargs)
245-
elif event == measurement.Event.START_STEP:
245+
elif event == measurement_base.Event.START_STEP:
246246
self._recorder.record_step_start_time(*args, **kwargs)
247-
elif event == measurement.Event.START_ACCELERATOR_INIT:
247+
elif event == measurement_base.Event.START_ACCELERATOR_INIT:
248248
self._recorder.record_tpu_init_start_time(*args, **kwargs)
249-
elif event == measurement.Event.END_ACCELERATOR_INIT:
249+
elif event == measurement_base.Event.END_ACCELERATOR_INIT:
250250
self._recorder.record_tpu_init_end_time(*args, **kwargs)
251-
elif event == measurement.Event.START_TRAINING_PREPARATION:
251+
elif event == measurement_base.Event.START_TRAINING_PREPARATION:
252252
self._recorder.record_training_preparation_start_time(*args, **kwargs)
253-
elif event == measurement.Event.END_TRAINING_PREPARATION:
253+
elif event == measurement_base.Event.END_TRAINING_PREPARATION:
254254
self._recorder.record_training_preparation_end_time(*args, **kwargs)
255-
elif event == measurement.Event.START_DATA_LOADING:
255+
elif event == measurement_base.Event.START_DATA_LOADING:
256256
self._recorder.record_data_loading_start_time(*args, **kwargs)
257-
elif event == measurement.Event.END_DATA_LOADING:
257+
elif event == measurement_base.Event.END_DATA_LOADING:
258258
self._recorder.record_data_loading_end_time(*args, **kwargs)
259-
elif event == measurement.Event.START_CUSTOM_BADPUT_EVENT:
259+
elif event == measurement_base.Event.START_CUSTOM_BADPUT_EVENT:
260260
self._recorder.record_custom_badput_event_start_time(*args, **kwargs)
261-
elif event == measurement.Event.END_CUSTOM_BADPUT_EVENT:
261+
elif event == measurement_base.Event.END_CUSTOM_BADPUT_EVENT:
262262
self._recorder.record_custom_badput_event_end_time(*args, **kwargs)
263263
else:
264264
logging.log_first_n(

axlearn/cloud/gcp/pathways_utils.py

Lines changed: 4 additions & 26 deletions
Original file line numberDiff line numberDiff line change
@@ -311,30 +311,6 @@ def _build_pathways_head_container(self) -> dict:
311311
}
312312
)
313313

314-
# pylint: disable=line-too-long
315-
env_list.append(
316-
{
317-
"name": "NUM_REPLICAS",
318-
"valueFrom": {
319-
"fieldRef": {
320-
"fieldPath": "metadata.annotations['jobset.sigs.k8s.io/replicatedjob-replicas']"
321-
}
322-
},
323-
}
324-
)
325-
# pylint: enable=line-too-long
326-
327-
env_list.append(
328-
{
329-
"name": "REPLICA_ID",
330-
"valueFrom": {
331-
"fieldRef": {
332-
"fieldPath": "metadata.annotations['jobset.sigs.k8s.io/job-index']"
333-
}
334-
},
335-
}
336-
)
337-
338314
head_container["env"] = env_list
339315

340316
cpu_req = f"{float(self.config.pathways_head_cpu) * 1000}m"
@@ -841,11 +817,13 @@ def default_config(cls):
841817
cfg = super().default_config()
842818
return cfg.set(inner=TPULeaderWorkerTemplate.default_config())
843819

844-
def __init__(self, cfg, *, bundler):
820+
def __init__(self, cfg: BaseLeaderWorkerTemplate.Config, *, bundler):
845821
super().__init__(cfg, bundler=bundler)
822+
cfg: PathwaysLeaderWorkerTemplate.Config = self.config
823+
846824
self._bundler = bundler
847825
self._inner: TPULeaderWorkerTemplate = cfg.inner.instantiate(bundler=self._bundler)
848-
self._tpu_type = infer_tpu_type(cfg.accelerator.instance_type)
826+
self._tpu_type = infer_tpu_type(cfg.inner.accelerator.instance_type)
849827
if self._tpu_type not in USER_FACING_NAME_TO_SYSTEM_CHARACTERISTICS:
850828
raise NotImplementedError(f"Missing system characteristics for {self._tpu_type}")
851829

axlearn/cloud/gcp/pathways_utils_test.py

Lines changed: 9 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -120,6 +120,15 @@ def test_build_pathways_head_pod(self, instance_type):
120120
}
121121
},
122122
)
123+
if env_pair["name"] == "REPLICA_ID":
124+
self.assertEqual(
125+
env_pair["valueFrom"],
126+
{
127+
"fieldRef": {
128+
"fieldPath": "metadata.annotations['jobset.sigs.k8s.io/job-index']"
129+
}
130+
},
131+
)
123132
if env_pair["name"] == "IFRT_PROXY_LARGE_TRANSFER_THRESHOLD":
124133
self.assertEqual(env_pair["value"], "1")
125134
if env_pair["name"] == "IFRT_PROXY_LARGE_TRANSFER_OPTIMIZATION_DIRECTORY":

axlearn/common/flash_attention/tpu_attention.py

Lines changed: 18 additions & 24 deletions
Original file line numberDiff line numberDiff line change
@@ -441,14 +441,12 @@ def kv_segment_ids_index_map(batch_index, head_index, q_seq_index, kv_seq_index)
441441
out_shape=out_shape,
442442
debug=debug,
443443
interpret=interpret,
444-
compiler_params=dict(
445-
mosaic=dict(
446-
dimension_semantics=(
447-
"parallel",
448-
"parallel",
449-
"parallel",
450-
"arbitrary",
451-
)
444+
compiler_params=pltpu.CompilerParams(
445+
dimension_semantics=(
446+
"parallel",
447+
"parallel",
448+
"parallel",
449+
"arbitrary",
452450
)
453451
),
454452
)(q, k, v, ab, q_segment_ids, kv_segment_ids)
@@ -649,14 +647,12 @@ def dkv_index_map(batch_index, head_index, kv_seq_index, _):
649647
out_shape=out_shapes,
650648
debug=debug,
651649
interpret=interpret,
652-
compiler_params=dict(
653-
mosaic=dict(
654-
dimension_semantics=(
655-
"parallel",
656-
"parallel",
657-
"parallel",
658-
"arbitrary",
659-
)
650+
compiler_params=pltpu.CompilerParams(
651+
dimension_semantics=(
652+
"parallel",
653+
"parallel",
654+
"parallel",
655+
"arbitrary",
660656
)
661657
),
662658
)(q, k, v, ab, q_segment_ids, kv_segment_ids, l, m, do, di)
@@ -842,14 +838,12 @@ def kv_segment_ids_index_map(batch_index, head_index, q_seq_index, kv_seq_index)
842838
out_shape=out_shapes,
843839
debug=debug,
844840
interpret=interpret,
845-
compiler_params=dict(
846-
mosaic=dict(
847-
dimension_semantics=(
848-
"parallel",
849-
"parallel",
850-
"parallel",
851-
"arbitrary",
852-
)
841+
compiler_params=pltpu.CompilerParams(
842+
dimension_semantics=(
843+
"parallel",
844+
"parallel",
845+
"parallel",
846+
"arbitrary",
853847
)
854848
),
855849
)(q, k, v, ab, q_segment_ids, kv_segment_ids, l, m, do, di)

axlearn/common/input_dispatch.py

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -117,7 +117,9 @@ class Config(BaseInputDispatcher.Config):
117117
def __init__(self, cfg: Config, *, parent: Optional[Module]):
118118
cfg = cfg.clone()
119119
cfg.num_physical_feeds = cfg.num_physical_feeds or jax.process_count()
120-
cfg.physical_feed_index = cfg.physical_feed_index or jax.process_index()
120+
cfg.physical_feed_index = (
121+
cfg.physical_feed_index if cfg.physical_feed_index is not None else jax.process_index()
122+
)
121123
if cfg.logical_feed_indices is None:
122124
num_logical_feeds = min(cfg.global_logical_batch_size, cfg.num_physical_feeds)
123125
cfg.logical_feed_indices = list(range(num_logical_feeds))

axlearn/common/input_tf_data_gcs_test.py

Lines changed: 5 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -76,7 +76,11 @@ def test_maybe_shard_examples(
7676
dataset_name=dataset_name,
7777
)
7878
if expected == "even split":
79-
shard_index = read_config.shard_index or jax.process_index()
79+
shard_index = (
80+
read_config.shard_index
81+
if read_config.shard_index is not None
82+
else jax.process_index()
83+
)
8084
expected_split = tfds.even_splits(split, n=required_shards, drop_remainder=False)[
8185
shard_index
8286
]

0 commit comments

Comments
 (0)