Skip to content

Commit 963383f

Browse files
committed
Merge branch 'main' into jax_0.8.0_py3.12_v7x
2 parents 41773c6 + 5a5b55d commit 963383f

175 files changed

Lines changed: 1391 additions & 3841 deletions

File tree

Some content is hidden

Large Commits have some content hidden by default. Use the searchbox below for content that may be hidden.

.github/workflows/stale.yml

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -15,3 +15,7 @@ jobs:
1515
close-pr-message: "This pull request was closed because it has been inactive for more than 7 days since being marked as stale. Please feel free to reopen it if you would like to continue."
1616
exempt-pr-labels: "ready-to-merge"
1717
stale-pr-label: "stale"
18+
permissions:
19+
actions: write
20+
issues: write
21+
pull-requests: write

Dockerfile

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -45,7 +45,7 @@ RUN pip install -qq --upgrade pip && \
4545
FROM base AS ci
4646

4747
# TODO(markblee): Remove gcp,vertexai_tensorboard from CI.
48-
RUN uv pip install -qq .[core,audio,orbax,dev,gcp,vertexai_tensorboard,open_api] && \
48+
RUN uv pip install -qq .[core,audio,orbax,dev,gcp,vertexai_tensorboard] && \
4949
uv cache clean
5050
COPY . .
5151

README.md

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,7 @@
11
# The AXLearn Library for Deep Learning
22

3+
[![build-and-test](https://github.com/apple/axlearn/actions/workflows/build.yml/badge.svg?branch=main)](https://github.com/apple/axlearn/actions/workflows/build.yml)
4+
35
**This library is under active development and the API is subject to change.**
46

57
## Table of Contents

axlearn/audio/encoder_asr.py

Lines changed: 15 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -16,7 +16,7 @@
1616
from axlearn.audio.frontend import LogMelFrontend
1717
from axlearn.audio.spectrum_augmenter import SpectrumAugmenter
1818
from axlearn.audio.subsamplers import ConvSubSampler
19-
from axlearn.common.attention import SinusoidalPositionalEmbedding
19+
from axlearn.common.attention import RepeatedTransformerLayer, SinusoidalPositionalEmbedding
2020
from axlearn.common.base_layer import BaseLayer
2121
from axlearn.common.config import REQUIRED, Required, config_class
2222
from axlearn.common.conformer import RepeatedConformerLayer
@@ -165,11 +165,23 @@ def forward(self, inputs: Tensor, *, paddings: Tensor) -> dict[str, Tensor]:
165165
- outputs: A Tensor of shape [batch_size, seq_len, output_dim].
166166
- output_paddings: A 0/1 Tensor of shape [batch_size, seq_len].
167167
"""
168+
cfg = self.config
168169
# [batch, seq_len, input_dim].
169170
x = self.input_linear(inputs)
170171
x = self.dropout(x)
171-
x = x + self.pos_emb(jnp.arange(x.shape[1]))
172-
x = self.context(inputs=x, paddings=paddings)
172+
173+
if isinstance(cfg.context, RepeatedConformerLayer.Config):
174+
x = x + self.pos_emb(jnp.arange(x.shape[1]))
175+
x = self.context(inputs=x, paddings=paddings)
176+
elif isinstance(cfg.context, RepeatedTransformerLayer.Config):
177+
# We don't need to do add pos_emb for transformer block
178+
x = self.context(data=x)
179+
x = x.data
180+
else:
181+
raise ValueError(
182+
f"The type of `self.context` ({cfg.context.klass}) "
183+
"is not supported by SpeechContextNetwork."
184+
)
173185
self._add_activation_summary(
174186
name="speech_context",
175187
activations=x,

axlearn/audio/encoder_asr_test.py

Lines changed: 84 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -10,6 +10,8 @@
1010
from axlearn.audio import frontend_utils
1111
from axlearn.audio.encoder_asr import ASREncoder, SpeechContextNetwork, SpeechFeatureLayer
1212
from axlearn.audio.test_utils import fake_audio
13+
from axlearn.common.attention import RepeatedTransformerLayer
14+
from axlearn.common.kv_cache.sliding_window_kv_cache import enable_sliding_window_attention
1315
from axlearn.common.module import functional as F
1416
from axlearn.common.test_utils import TestCase
1517
from axlearn.common.utils import Tensor, shapes
@@ -173,6 +175,88 @@ def test_speech_context_network(self, is_training: bool):
173175
output_collections.summaries["activations/speech_context_norm"].weight, weights
174176
)
175177

178+
@parameterized.parameters([True, False])
179+
def test_transformer(self, is_training: bool) -> None:
180+
"""Test the code branch with RepeatedTransformerLayer as context layer.
181+
182+
Args:
183+
is_training: Whether the is_training code path is tested.
184+
"""
185+
input_dim, output_dim, dropout_rate, num_layers = 32, 16, 0.2, 2
186+
num_heads = 8
187+
hidden_dim = 4 * input_dim
188+
189+
cfg = SpeechContextNetwork.default_config().set(
190+
input_dim=input_dim, output_dim=output_dim, dtype=jnp.float64
191+
)
192+
cfg.dropout.rate = dropout_rate
193+
cfg.context = RepeatedTransformerLayer.default_config().set(num_layers=num_layers)
194+
attention = cfg.context.layer.self_attention.attention
195+
attention.num_heads = num_heads
196+
attention = enable_sliding_window_attention(attention, sliding_window_size=3)
197+
cfg.context.layer.self_attention.attention = attention
198+
# Dropout in transformer
199+
cfg.context.layer.self_attention.dropout.rate = dropout_rate
200+
cfg.context.layer.feed_forward.set(
201+
hidden_dim=hidden_dim,
202+
)
203+
204+
# Initialize layer parameters.
205+
prng_key = jax.random.PRNGKey(123)
206+
prng_key, init_key, input_key, length_key = jax.random.split(prng_key, num=4)
207+
layer = cfg.set(name="test").instantiate(parent=None)
208+
layer_params = layer.initialize_parameters_recursively(init_key)
209+
210+
# Generate inputs.
211+
batch_size, seq_len = 4, 10
212+
inputs = jnp.tile(
213+
jax.random.normal(input_key, [batch_size // 2, seq_len, input_dim]), [2, 1, 1]
214+
)
215+
lengths = jnp.tile(
216+
jax.random.randint(length_key, shape=[batch_size // 2, 1], minval=0, maxval=seq_len),
217+
[2, 1],
218+
)
219+
paddings = jnp.arange(seq_len)[None, :] >= lengths
220+
padding_data = jax.random.normal(jax.random.PRNGKey(135), inputs.shape)
221+
inputs = jnp.where(paddings[..., None], padding_data, inputs)
222+
223+
# Compute outputs.
224+
output_batch, output_collections = F(
225+
layer,
226+
inputs=dict(inputs=inputs, paddings=paddings),
227+
is_training=is_training,
228+
prng_key=prng_key,
229+
state=layer_params,
230+
)
231+
outputs, output_paddings = output_batch["outputs"], output_batch["paddings"]
232+
self.assertSequenceEqual(outputs.shape, (batch_size, seq_len, output_dim))
233+
self.assertTrue(jnp.all(output_paddings == paddings))
234+
235+
# If is_training, outputs should always be different due to augmentation.
236+
# Otherwise, outputs should be the same despite differences in padding.
237+
self.assertEqual(not is_training, bool(jnp.allclose(outputs[:2], outputs[2:])))
238+
239+
outputs = outputs * (1 - output_paddings[:, :, None])
240+
weights = jnp.sum(1 - output_paddings)
241+
output_norms = jnp.sqrt(jnp.sum(outputs**2, axis=2)) / jnp.sqrt(output_dim)
242+
expected_outputs_mean = jnp.sum(outputs) / weights / output_dim
243+
expected_outputs_norm = jnp.sum(output_norms) / weights
244+
245+
self.assertNestedAllClose(
246+
output_collections.summaries["activations/speech_context_mean"].mean,
247+
expected_outputs_mean,
248+
)
249+
self.assertNestedAllClose(
250+
output_collections.summaries["activations/speech_context_norm"].mean,
251+
expected_outputs_norm,
252+
)
253+
self.assertNestedAllClose(
254+
output_collections.summaries["activations/speech_context_mean"].weight, weights
255+
)
256+
self.assertNestedAllClose(
257+
output_collections.summaries["activations/speech_context_norm"].weight, weights
258+
)
259+
176260

177261
class ASREncoderTest(TestCase):
178262
"""Tests ASREncoder."""

axlearn/cloud/gcp/jobset_utils.py

Lines changed: 23 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -502,6 +502,29 @@ def _build_container(self) -> Nested[Any]:
502502
k8s_env_vars.append(
503503
{"name": "NODE_NAME", "valueFrom": {"fieldRef": {"fieldPath": "spec.nodeName"}}}
504504
)
505+
# pylint: disable=line-too-long
506+
k8s_env_vars.append(
507+
{
508+
"name": "NUM_REPLICAS",
509+
"valueFrom": {
510+
"fieldRef": {
511+
"fieldPath": "metadata.annotations['jobset.sigs.k8s.io/replicatedjob-replicas']"
512+
}
513+
},
514+
}
515+
)
516+
# pylint: enable=line-too-long
517+
518+
k8s_env_vars.append(
519+
{
520+
"name": "REPLICA_ID",
521+
"valueFrom": {
522+
"fieldRef": {
523+
"fieldPath": "metadata.annotations['jobset.sigs.k8s.io/job-index']"
524+
}
525+
},
526+
}
527+
)
505528

506529
return dict(
507530
name=cfg.name,

axlearn/cloud/gcp/jobset_utils_test.py

Lines changed: 12 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -359,6 +359,18 @@ def test_build_pod(
359359
container_env["NODE_IP"]["valueFrom"]["fieldRef"]["fieldPath"],
360360
)
361361

362+
# Verify NUM_REPLICAS in container env.
363+
self.assertEqual(
364+
"metadata.annotations['jobset.sigs.k8s.io/replicatedjob-replicas']",
365+
container_env["NUM_REPLICAS"]["valueFrom"]["fieldRef"]["fieldPath"],
366+
)
367+
368+
# Verify REPLICA_ID in container env.
369+
self.assertEqual(
370+
"metadata.annotations['jobset.sigs.k8s.io/job-index']",
371+
container_env["REPLICA_ID"]["valueFrom"]["fieldRef"]["fieldPath"],
372+
)
373+
362374
# Verify uploader container specs
363375
self.assertEqual(len(pod_spec["initContainers"]), 1)
364376

axlearn/common/array_serialization_test.py

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -466,6 +466,7 @@ def test_shard_info_partially_replicated(
466466
devices = mesh_utils.create_device_mesh((8,))
467467
mesh = jax.sharding.Mesh(devices.reshape((4, 2)), ("x", "y"))
468468
sharding = jax.sharding.NamedSharding(mesh, jax.sharding.PartitionSpec(None, "y"))
469+
469470
arr = jax.device_put(single_device_arr, sharding)
470471

471472
replica_count = _num_replicas_per_shard(arr)
@@ -486,6 +487,7 @@ def test_shard_info_fully_sharded(self, max_data_shard_degree: int, shard_thresh
486487
devices = mesh_utils.create_device_mesh((8,))
487488
mesh = jax.sharding.Mesh(devices.reshape((4, 2)), ("x", "y"))
488489
sharding = jax.sharding.NamedSharding(mesh, jax.sharding.PartitionSpec("x", "y"))
490+
489491
arr = jax.device_put(single_device_arr, sharding)
490492

491493
replica_count = _num_replicas_per_shard(arr)
@@ -509,6 +511,7 @@ def test_shard_info_fully_replicated(
509511
devices = mesh_utils.create_device_mesh((8,))
510512
mesh = jax.sharding.Mesh(devices, "x")
511513
sharding = jax.sharding.NamedSharding(mesh, jax.sharding.PartitionSpec(None))
514+
512515
arr = jax.device_put(single_device_arr, sharding)
513516

514517
replica_count = _num_replicas_per_shard(arr)

axlearn/common/causal_lm_test.py

Lines changed: 7 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -8,13 +8,12 @@
88
import jax
99
import jax.random
1010
import numpy as np
11-
import pytest
1211
from absl.testing import absltest, parameterized
1312
from jax import numpy as jnp
1413
from jax.experimental.pjit import pjit
1514
from transformers.models.gpt2 import modeling_gpt2 as hf_gpt2
1615

17-
from axlearn.common import causal_lm, utils
16+
from axlearn.common import causal_lm
1817
from axlearn.common.attention import (
1918
BaseStackedTransformerLayer,
2019
CausalAttentionLogitBiasLayer,
@@ -476,14 +475,12 @@ def test_no_conflict(
476475
)
477476

478477
# TODO(markblee): Add a pytest marker for multi-device tests.
479-
@pytest.mark.skipif(
480-
jax.device_count() != 4 or jax.process_count() != 1,
481-
reason=(
482-
"Incorrect device & process count for mesh.\n"
483-
"Use XLA_FLAGS=--xla_force_host_platform_device_count=4 to run locally."
484-
),
485-
)
486478
def test_constrain_input_batch(self):
479+
if jax.device_count() != 4 or jax.process_count() != 1:
480+
self.skipTest(
481+
"Incorrect device & process count for mesh.\n"
482+
"Use XLA_FLAGS=--xla_force_host_platform_device_count=4 to run locally."
483+
)
487484
model = (
488485
self._model_config(vocab_size=10, seq_len=10)
489486
.set(
@@ -817,5 +814,4 @@ def loss_fn(model_params, inputs):
817814

818815

819816
if __name__ == "__main__":
820-
with utils.numeric_checks(True):
821-
absltest.main()
817+
absltest.main()

axlearn/common/checkpointer_test.py

Lines changed: 8 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -21,6 +21,7 @@
2121

2222
import jax
2323
import orbax.checkpoint as ocp
24+
import pytest
2425
import tensorflow as tf
2526
from absl import logging
2627
from absl.testing import absltest, parameterized
@@ -171,6 +172,7 @@ def test_save_and_restore(self, checkpointer_cls: Type[BaseCheckpointer]):
171172
ckpt.stop()
172173

173174
@parameterized.parameters(Checkpointer, OrbaxCheckpointer)
175+
@pytest.mark.for_8_devices
174176
def test_save_and_restore_mesh(self, checkpointer_cls: Type[BaseCheckpointer]):
175177
"""Tests that we can save with one sharding and restore with a different sharding."""
176178
mesh_shape = (4, 2)
@@ -245,6 +247,7 @@ def state_specs(state, partition_spec):
245247
num_files=6, # 1 array 4 shards (2 model, 2 data) + 1 array 2 shards (small array).
246248
),
247249
)
250+
@pytest.mark.for_8_devices
248251
def test_save_restore_files_count(
249252
self, max_data_shard_degree: int, shard_threshold_bytes: int, num_files: int
250253
):
@@ -668,9 +671,11 @@ def patch_tf_io_behavior(*args):
668671
# pylint: disable=line-too-long
669672
with (
670673
_mesh(mesh_shape),
671-
mock.patch("axlearn.common.file_system.listdir", patch_tf_io_behavior)
672-
if listdir_add_trailing_slash
673-
else nullcontext(),
674+
(
675+
mock.patch("axlearn.common.file_system.listdir", patch_tf_io_behavior)
676+
if listdir_add_trailing_slash
677+
else nullcontext()
678+
),
674679
tempfile.TemporaryDirectory() as temp_dir,
675680
):
676681
cfg = Checkpointer.default_config().set(

0 commit comments

Comments
 (0)