Code
Hub
Workspaces
Following
Trending
Connect
MCP
copy
Create free account
hub
/
github.com/apple/axlearn
/ endpoints
Endpoints
180 in github.com/apple/axlearn
⨍
Functions
7,799
◇
Types & classes
2,118
↳
Endpoints
180
Route
_gauge
functools.singledispatchmethod
axlearn/cloud/gcp/monitoring/tpu_client_test.py:None
Route
scan_fn
functools.partial( jax.ad_checkpoint.checkpoint, prevent_cse=False, policy
axlearn/common/pipeline.py:None
Route
test_activations
parameterized.parameters( dict(activation=None, expected_activation=[None, None], output_dim=3),
axlearn/audio/subsamplers_test.py:None
Route
test_active_command_malformed
mock.patch("subprocess.Popen")
axlearn/cloud/common/bastion_test.py:None
Route
test_add_prefix_concat_sequence_pair
parameterized.parameters( # Test a simple classification example. dict( input_exam
axlearn/common/input_glue_test.py:None
Route
test_add_token_type_ids_multiple_examples
parameterized.parameters( # For ragged tensors, for each of the three example, # query has 8 t
axlearn/common/input_text_test.py:None
Route
test_add_token_type_ids_single_example
parameterized.parameters( # For ragged tensors, query has 8 tokens, answer has 5 tokens, and url has 6
axlearn/common/input_text_test.py:None
Route
test_add_value_rms_norm_summary
parameterized.parameters( dict(rms_norm_summary=[]), dict(rms_norm_summary=["linear2_outputs"]
axlearn/common/attention_test.py:None
Route
test_apply_t5_mask
parameterized.parameters( # Test an empty input. dict( source_ids=tf.constant([]),
axlearn/common/input_t5_test.py:None
Route
test_asr_input
parameterized.parameters( # Not dropping any. dict(max_source_len=5, max_target_len=6, expect_
axlearn/experiments/audio/conformer/common_test.py:None
Route
test_asr_input
parameterized.parameters( dict( # Test when truncate=False. max_speech_len=5,
axlearn/audio/input_asr_test.py:None
Route
test_augment_text_from_inputs_targets_pretokenized
parameterized.parameters( dict( replace_newlines_with="<n>", source_key="input
axlearn/common/input_lm_test.py:None
Route
test_backward
parameterized.product(normalize_inputs=(True, False), normalize_codebook=(True, False))
axlearn/common/quantizer_test.py:None
Route
test_batch_norm
parameterized.parameters( [ dict(inputs_shape=[2, 3, 6], segment_ids=None), di
axlearn/common/layers_test.py:None
Route
test_build_and_push
parameterized.product( platform=[None, "test-platform"], target=[None, "test-target"], )
axlearn/cloud/common/bundler_test.py:None
Route
test_build_pod
parameterized.product( [ dict( env={}, reservation=None,
axlearn/cloud/gcp/jobset_utils_test.py:None
Route
test_bwd_against_ref
parameterized.product( [ dict(batch_size=1, num_heads=1, seq_len=2048, per_head_dim=64),
axlearn/common/flash_attention/neuron_attention_test.py:None
Route
test_call_unclean
mock.patch( f"{git_summary_path}.to_labels", return_value={"fake-kabel": "fake-label-value"},
axlearn/cloud/common/bundler_test.py:None
Route
test_causal
parameterized.product( base_cfg=( attention.MultiheadAttention.default_config(),
axlearn/common/attention_test.py:None
Route
test_checkpoint_policy
parameterized.product( save_input_iterator=[False, True], restore_input_iterator=[False, True]
axlearn/common/trainer_test.py:None
Route
test_classification_metrics
parameterized.parameters( # Test case for multiple precision, recall level { "inpu
axlearn/common/eval_classification_test.py:None
Route
test_composite_attention_has_bias
parameterized.parameters( [ attention_bias.CompositeAttentionBias( [attent
axlearn/common/attention_bias_test.py:None
Route
test_compute_fan_axes
parameterized.parameters( ( MultiheadInputLinear, FanAxes(in_axis=0, out_axis=
axlearn/common/attention_test.py:None
Route
test_connect_amqp_connection_error
mock.patch.dict(os.environ, {"RABBITMQ_USER": "test_user", "RABBITMQ_PASSWORD": "test_pass"})
axlearn/cloud/common/event_queue_test.py:None
Route
test_connect_missing_credentials
mock.patch.dict(os.environ, {"RABBITMQ_USER": "", "RABBITMQ_PASSWORD": ""})
axlearn/cloud/common/event_queue_test.py:None
Route
test_connect_success
mock.patch.dict(os.environ, {"RABBITMQ_USER": "test_user", "RABBITMQ_PASSWORD": "test_pass"})
axlearn/cloud/common/event_queue_test.py:None
Route
test_conv2d
parameterized.named_parameters( { "testcase_name": "1x1", "window": (1, 1),
axlearn/common/convolution_test.py:None
Route
test_conv3d
parameterized.named_parameters( { "testcase_name": "1x1x1", "window": (1, 1, 1
axlearn/common/convolution_test.py:None
Route
test_convert_to_monitored_layer_config
parameterized.parameters(None, compute_grad_percentile_no_clip_fn, top_k_clip_fn)
axlearn/common/gradient_monitor_test.py:None
Route
test_create_node_pool_fire_and_forget
parameterized.parameters( dict(fire_and_forget=True, exception=None, expected=None), dict(fire
axlearn/cloud/gcp/node_pool_test.py:None
Route
test_decode_against_ref
parameterized.product( [ dict(zip(["batch_size", "seq_len", "num_heads", "per_head_dim"],
axlearn/common/flash_attention/decoding_test.py:None
Route
test_decoding
parameterized.product( _TEST_CONFIGS, backend=["cpu", "gpu", "tpu"], bias_type=["causa
axlearn/common/flash_attention/utils_test.py:None
Route
test_default_input_dispatcher
parameterized.parameters( {"target_pid": 0}, {"target_pid": 1}, {"target_pid": 2}, {"target_pid": 3}
axlearn/common/elastic_input_test.py:None
Route
test_delete_node_pool_fire_and_forget
parameterized.parameters( dict(fire_and_forget=True, exception=None, expected=None), dict(fire
axlearn/cloud/gcp/node_pool_test.py:None
Route
test_delete_node_pools
parameterized.parameters( dict(names=["node_pool0"], wait_timeout=0, expected_delete_call_count=1),
axlearn/cloud/gcp/node_pool_test.py:None
Route
test_dependencies
parameterized.parameters( # dependencies defines a mapping (src, dst, dst_key). If a calculator is lis
axlearn/common/evaler_test.py:None
Route
test_dit_attn
parameterized.parameters(["prenorm", "postnorm", "hybridnorm"])
axlearn/common/dit_test.py:None
Route
test_dit_ffn
parameterized.parameters(["prenorm", "postnorm", "hybridnorm"])
axlearn/common/dit_test.py:None
Route
test_einsum_maybe_quantized
parameterized.product( b=[2, 16], d=[4, 32], h=[8, 64], quantization_type_and_
axlearn/common/quantized_dot_general/layers_test.py:None
Route
test_element_spec
parameterized.parameters( dict( dispatcher=SpmdInputDispatcher, # No change wi
axlearn/common/input_base_test.py:None
Route
test_ema_params_converter
parameterized.parameters(["with_target_ema", "with_learner_no_ema", "with_no_learner"])
axlearn/common/state_builder_test.py:None
Route
test_embed_partition_specs_constraint
mock.patch("axlearn.common.utils.with_sharding_constraint")
axlearn/common/layers_test.py:None
Route
test_en_normalizer
parameterized.named_parameters( {"testcase_name": "answer_n/a", "answer": "n/a", "expected": "n"},
axlearn/vision/metrics_vqa_test.py:None
Route
test_enable_monitoring_enabled_by_default
mock.patch("jax.process_index", return_value=0)
axlearn/cloud/gcp/measurement_test.py:None
Route
test_enable_monitoring_explicitly_disabled
mock.patch("jax.process_index", return_value=0)
axlearn/cloud/gcp/measurement_test.py:None
Route
test_enable_monitoring_explicitly_enabled
mock.patch("jax.process_index", return_value=0)
axlearn/cloud/gcp/measurement_test.py:None
Route
test_extend_step
parameterized.parameters( dict( # Baseline use_cross_attention=False, stack_c
axlearn/common/decoder_test.py:None
Route
test_extend_step
parameterized.product( dtype=(jnp.float32, jnp.float16, jnp.bfloat16), per_dim_scale=(None, Pe
axlearn/common/attention_test.py:None
Route
test_fake_text2text_lm_input
parameterized.parameters( # Encoder-decoder. { "is_training": False, "
axlearn/common/input_lm_test.py:None
Route
test_fake_text_lm_training_data
parameterized.parameters( dict( packing_method=PackingMethodType.EOS_DELIM_MASK,
axlearn/common/input_lm_test.py:None
Route
test_fake_text_lm_training_data
parameterized.parameters( dict( expected_batches=[ { "
axlearn/common/input_grain_lm_test.py:None
Route
test_filter_by_length
parameterized.parameters( dict( input_key="inputs", inputs=[{"inputs": []}, {"
axlearn/audio/input_asr_test.py:None
Route
test_filter_module_outputs
parameterized.parameters( # Tests basic key lookup. dict( expected={"target_labels
axlearn/common/loss_metrics_test.py:None
Route
test_forward
parameterized.product( # Parameterize how source padding is represented: # 1. none: Test no pa
axlearn/common/encoder_decoder_test.py:None
Route
test_forward
parameterized.product(num_groups=(1, 2), input_mean=(0.0, -0.5))
axlearn/common/quantizer_test.py:None
Route
test_forward
parameterized.product( _TEST_CONFIGS, backend=["cpu", "gpu", "tpu"], bias_type=["full"
axlearn/common/flash_attention/utils_test.py:None
Route
test_forward_with_normalization
parameterized.product(norm_inputs=(False, True), norm_codebook=(False, True))
axlearn/common/quantizer_test.py:None
Route
test_from_flags
parameterized.product( name=[None, "test-name"], output_dir=[None, "test-output"], )
axlearn/cloud/gcp/jobs/cpu_runner_test.py:None
Route
test_full_partition
mock.patch("axlearn.common.utils.input_partition_spec")
axlearn/common/utils_test.py:None
Route
test_fwd_against_ref
parameterized.product( [ dict(batch_size=1, seq_len=2048, num_heads=1, per_head_dim=64),
axlearn/common/flash_attention/neuron_attention_test.py:None
Route
test_get_cloud_build_status_correctly_sets_last_known_region_for_build
mock.patch("axlearn.cloud.gcp.cloud_build._get_latest_build_status_in_region")
axlearn/cloud/gcp/cloud_build_test.py:None
Route
test_get_cloud_build_status_correctly_uses_preset_last_known_region_for_build
mock.patch("axlearn.cloud.gcp.cloud_build._get_latest_build_status_in_region")
axlearn/cloud/gcp/cloud_build_test.py:None
Route
test_get_cloud_build_status_failure_with_two_regions
mock.patch("axlearn.cloud.gcp.cloud_build._get_latest_build_status_in_region")
axlearn/cloud/gcp/cloud_build_test.py:None
Route
test_get_cloud_build_status_returns_none_with_no_regions
mock.patch("axlearn.cloud.gcp.cloud_build._get_latest_build_status_in_region")
axlearn/cloud/gcp/cloud_build_test.py:None
Route
test_get_cloud_build_status_success_with_only_global_region
mock.patch("axlearn.cloud.gcp.cloud_build._get_latest_build_status_in_region")
axlearn/cloud/gcp/cloud_build_test.py:None
Route
test_get_cloud_build_status_success_with_two_regions
mock.patch("axlearn.cloud.gcp.cloud_build._get_latest_build_status_in_region")
axlearn/cloud/gcp/cloud_build_test.py:None
Route
test_get_flink_job_status
mock.patch("axlearn.cloud.gcp.runners.gke.requests.get")
axlearn/cloud/gcp/runners/gke_test.py:None
Route
test_get_flink_job_status_raises_on_request_error
mock.patch("axlearn.cloud.gcp.runners.gke.requests.get")
axlearn/cloud/gcp/runners/gke_test.py:None
Route
test_get_latest_build_status_in_region_raises_exception
mock.patch("axlearn.cloud.gcp.cloud_build.cloudbuild_v1.CloudBuildClient")
axlearn/cloud/gcp/cloud_build_test.py:None
Route
test_get_latest_build_status_in_region_returns_failure_status
mock.patch("axlearn.cloud.gcp.cloud_build._list_builds_in_region")
axlearn/cloud/gcp/cloud_build_test.py:None
Route
test_get_latest_build_status_in_region_returns_none_with_empty_builds
mock.patch("axlearn.cloud.gcp.cloud_build._list_builds_in_region")
axlearn/cloud/gcp/cloud_build_test.py:None
Route
test_get_latest_build_status_in_region_returns_status_unknown
mock.patch("axlearn.cloud.gcp.cloud_build._list_builds_in_region")
axlearn/cloud/gcp/cloud_build_test.py:None
Route
test_get_latest_build_status_in_region_returns_success_status
mock.patch("axlearn.cloud.gcp.cloud_build._list_builds_in_region")
axlearn/cloud/gcp/cloud_build_test.py:None
Route
test_get_latest_build_status_in_region_with_one_region_success
mock.patch("axlearn.cloud.gcp.cloud_build._list_builds_in_region")
axlearn/cloud/gcp/cloud_build_test.py:None
Route
test_get_or_prompt_project
parameterized.parameters( # Test basic case. dict(inputs=[], labels=["b"], expected="project0"
axlearn/cloud/common/config_test.py:None
Route
test_get_status
parameterized.product( config=( # Conditions is set, so we use it. GetStatusTe
axlearn/cloud/gcp/runners/gke_test.py:None
Route
test_gke_gateway_route_and_telemetry_combination
parameterized.parameters( (False, False), # Neither flag enabled (True, False), # Only gke_g
axlearn/cloud/gcp/pathways_utils_test.py:None
Route
test_gqa_extend_step
parameterized.product( dtype=(jnp.float32, jnp.float16, jnp.bfloat16), per_dim_scale=(None, Pe
axlearn/common/attention_test.py:None
Route
test_gqa_forward
parameterized.product( dtype=(jnp.float32, jnp.float16, jnp.bfloat16), per_dim_scale=(None, Pe
axlearn/common/attention_test.py:None
Route
test_gqa_prefill_states
parameterized.product( dtype=(jnp.float32, jnp.float16, jnp.bfloat16), per_dim_scale=(None, Pe
axlearn/common/attention_test.py:None
Route
test_gradient_clipping_implementation
parameterized.product( x=[ jax.random.normal(jax.random.PRNGKey(0), (2, 8, 64)),
axlearn/common/gradient_monitor_test.py:None
Route
test_group_norm
parameterized.parameters( [ dict(inputs_shape=[2, 3, 6]), dict(inputs_shape=[2
axlearn/common/layers_test.py:None
Route
test_has_bias
parameterized.parameters( [attention_bias.ZeroAttentionBias(), False], [ attention
axlearn/common/attention_bias_test.py:None
Route
test_incremental_prefill
parameterized.product( _TEST_CONFIGS, backend=["cpu", "tpu"], bias_type=["causal", "sl
axlearn/common/flash_attention/utils_test.py:None
Route
test_input_dispatcher
parameterized.parameters( # In the most common use cases, users only specify `global_logical_batch_siz
axlearn/common/input_dispatch_test.py:None
Route
test_is_kv_sharing
parameterized.parameters( dict(cfg=QKVLinear.default_config(), expected=False), dict(cfg=Fused
axlearn/common/attention_test.py:None
Route
test_is_kv_sharing
parameterized.parameters( dict(cfg=LoraFusedQKVLinear.default_config(), expected=False), dict(
axlearn/common/lora_test.py:None
Route
test_layer_structure
parameterized.parameters( "prenorm", "postnorm", "hybridnorm", "nonorm", )
axlearn/common/mixture_of_experts_test.py:None
Route
test_list_available_regions_raises_api_error
mock.patch("axlearn.cloud.gcp.cloud_build.logging.error")
axlearn/cloud/gcp/cloud_build_test.py:None
Route
test_list_available_regions_success
mock.patch("axlearn.cloud.gcp.cloud_build.RegionsClient")
axlearn/cloud/gcp/cloud_build_test.py:None
Route
test_list_builds_in_region_failure_with_list_builds_exception
mock.patch("axlearn.cloud.gcp.cloud_build.cloudbuild_v1.CloudBuildClient")
axlearn/cloud/gcp/cloud_build_test.py:None
Route
test_list_builds_in_region_success_with_empty_builds
mock.patch("axlearn.cloud.gcp.cloud_build.cloudbuild_v1.CloudBuildClient")
axlearn/cloud/gcp/cloud_build_test.py:None
Route
test_list_builds_in_region_success_with_global_region
mock.patch("axlearn.cloud.gcp.cloud_build.cloudbuild_v1.CloudBuildClient")
axlearn/cloud/gcp/cloud_build_test.py:None
Route
test_list_builds_in_region_success_with_non_global_region
mock.patch("axlearn.cloud.gcp.cloud_build.cloudbuild_v1.CloudBuildClient")
axlearn/cloud/gcp/cloud_build_test.py:None
Route
test_lm_from_seq2seq_text_preprocessor
parameterized.parameters( { "input_data_type": InputDataType.SEQ2SEQ_MASK, "ma
axlearn/common/input_lm_test.py:None
Route
test_log_mel_spectrogram
parameterized.product( # Inputs are [batch, num_frames, frame_size]. input_shape=[(5, 1298, 40
axlearn/audio/frontend_utils_test.py:None
Route
test_magnitude_spectrogram
parameterized.product( # Inputs are [batch, num_frames, frame_size]. input_shape=[(5, 1298, 40
axlearn/audio/frontend_utils_test.py:None
Route
test_make_autoregressive_checkpointing
parameterized.parameters( dict( target_labels=[ np.array([33, 33]),
axlearn/common/input_grain_lm_test.py:None
Route
test_make_autoregressive_inputs
parameterized.parameters( # Test a case without windowing. dict( target_labels=[
axlearn/common/input_grain_lm_test.py:None
Route
test_make_autoregressive_inputs_passthrough_keys
parameterized.parameters( { "expected_inputs": { "prefix": [0],
axlearn/common/input_lm_test.py:None
next →
1–100 of 180, ranked by callers