MCPcopy Create free account

hub / github.com/apple/axlearn / functions

Functions7,799 in github.com/apple/axlearn

↓ 5 callersFunctiontokenize
Tokenizes output_features using SeqIO. Args: output_features: A mapping from field name to seqio.preprocessors.OutputFeaturesType.
axlearn/common/input_text.py:134
↓ 5 callersFunctiontop_k_logits
Build a function that returns logits suitably normalized for top-k sampling. The returned function does many reductions over the last axis of the
axlearn/common/logit_modifiers.py:124
↓ 5 callersFunctionunflatten_decoding_dim
Unflattens the first, flat batch*decoding dimension of a non-scalar array.
axlearn/common/decoding.py:67
↓ 5 callersFunctionunwrap
Unwraps an image produced by wrap. Where there is a 0 in the last channel for every spatial position, the rest of the three channels in that
axlearn/vision/augment.py:565
↓ 5 callersMethodupdate
Update value at given index and propagate changes upward.
axlearn/common/segment_tree.py:61
↓ 5 callersFunctionvalidate_jobset_name
Validates JobSet name (e.g. TPUs, VMs, jobs) to ensure compat with GKE. Raises: ValueError: If name is invalid.
axlearn/cloud/gcp/utils.py:110
↓ 5 callersMethodvlog_is_on
(self, level: int)
axlearn/common/module.py:909
↓ 5 callersFunctionwrap
Returns 'image' with an extra channel set to all 1s.
axlearn/vision/augment.py:557
↓ 4 callersMethod__init__
(self, cfg: Config, *, parent: Optional[Module])
axlearn/vision/coca.py:843
↓ 4 callersMethod__init__
(self, cfg: Config, *, parent: Module)
axlearn/common/layers.py:199
↓ 4 callersMethod__init__
(self, cfg: Config, *, parent=None)
axlearn/common/trainer_test.py:94
↓ 4 callersMethod__init__
(self, cfg: Config, *, parent: Optional[Module])
axlearn/common/t5.py:266
↓ 4 callersMethod__init__
( self, cfg: Config, *, parent: Optional[Module], model: BaseModel,
axlearn/common/evaler.py:586
↓ 4 callersMethod__init__
(self, cfg: Config, *, parent: Module)
axlearn/common/mixture_of_experts.py:489
↓ 4 callersMethod_broadcast_value
Broadcasts `value` to a canonical 4 dimensional attention bias shape. Raises: ValueError: If the shape of `value` is not 2, 3, or
axlearn/common/attention_bias.py:124
↓ 4 callersMethod_build_backend_policy
Builds a config for a GCPBackendPolicy. Returns: A nested dict corresponding to a K8s GCPBackendPolicy config.
axlearn/cloud/gcp/k8s_backend_policy.py:52
↓ 4 callersMethod_build_jobset
Builds a config for a JobSet, which is a set of Jobs. https://github.com/kubernetes-sigs/jobset/blob/d49514bee57da8ac9aec2fcea06c3a13c21afeae
axlearn/cloud/gcp/job.py:217
↓ 4 callersMethod_build_pathways_head_pod
Builds a pathways head pod. The pod includes a head container, a proxy container and a resource manager container.
axlearn/cloud/gcp/pathways_utils.py:666
↓ 4 callersMethod_byte_encode_tf
Applies bytes_to_unicode mapping in tf. TODO(markblee): Consider baking this into SentencePiece normalization rules. Reference:
axlearn/common/vocabulary_bpe.py:266
↓ 4 callersMethod_cache_dtype
(self, dtype: jnp.dtype)
axlearn/common/kv_cache/base_kv_cache.py:64
↓ 4 callersMethod_call_model
Computes self._model.method(input_batch). Should be called inside pjit(). Args: method: The model method. pr
axlearn/common/evaler.py:205
↓ 4 callersFunction_common_flags
(job: BaseBastionManagedJob.Config)
axlearn/cloud/gcp/jobs/launch_test.py:141
↓ 4 callersMethod_compare_layers
( self, *stack_configs, dtype=jnp.float32, remat_spec=None, batch_size
axlearn/common/attention_test.py:5240
↓ 4 callersFunction_compare_slice_selection
Compares expected and deployed slice-selection annotations. Both None means no topology is configured, which is considered equal. Otherwise p
axlearn/cloud/gcp/runners/gke.py:140
↓ 4 callersFunction_compute_area
(y1: Tensor, x1: Tensor, y2: Tensor, x2: Tensor)
axlearn/common/loss.py:969
↓ 4 callersMethod_compute_context
Compute attention context. Args: probs: probs tensor, [batch, num_heads, target_length, source_length]. v_proj: value
axlearn/common/attention.py:2220
↓ 4 callersMethod_compute_fan_axes
(self, name: str, parameter_spec: ParameterSpec)
axlearn/common/base_layer_test.py:587
↓ 4 callersMethod_compute_logits
Compute attention logits. Args: q_proj: query tensor, [batch, target_length, num_heads, per_head_dim]. k_proj: key te
axlearn/common/attention.py:2208
↓ 4 callersMethod_delete
(self)
axlearn/cloud/gcp/jobs/dataflow.py:257
↓ 4 callersMethod_execute
(self)
axlearn/cloud/gcp/jobs/dataflow.py:263
↓ 4 callersFunction_flag_values_from_dict
(flag_values: dict)
axlearn/common/launch_trainer_test.py:29
↓ 4 callersFunction_gather_beams
Gathers the beam slices indexed by beam_indices into new beam array. Args: nested: NestedTensor or scalars (the latter ignored).
axlearn/common/decoding.py:78
↓ 4 callersFunction_gcp_settings_from_active_config
(key: str)
axlearn/cloud/gcp/config.py:57
↓ 4 callersMethod_get_flink_cluster_name
(self)
axlearn/cloud/gcp/job_flink.py:185
↓ 4 callersMethod_get_notary_configmap_name
Returns the appropriate ConfigMap name based on telemetry setting. Args: config_type: Either "http" or "grpc" Returns:
axlearn/cloud/gcp/pathways_utils.py:1324
↓ 4 callersMethod_get_num_of_tpu_nodes
(self, system: _SystemCharacteristics)
axlearn/cloud/gcp/job_flink.py:203
↓ 4 callersMethod_get_status
Gets current job status.
axlearn/cloud/gcp/jobs/cpu_runner.py:258
↓ 4 callersFunction_infer_num_partitions
Returns the number of partitions along each dim.
axlearn/common/host_array_test.py:35
↓ 4 callersMethod_input_config
( self, numbers: list[int], *, batch_size=2, repeat=1, out_sig
axlearn/common/input_composite_test.py:41
↓ 4 callersFunction_int32_binary_search
Binary search to find the largest finite int32 value for which the predicate is False. Ref: <https://github.com/google-research/t5x/blob/79998013
axlearn/common/logit_modifiers.py:220
↓ 4 callersFunction_is_supported
(*, platform: str, mesh_shape: MeshShape)
axlearn/common/host_array_test.py:58
↓ 4 callersMethod_job_config
(self, *, command: str, **kwargs)
axlearn/cloud/gcp/runners/gke_test.py:157
↓ 4 callersMethod_learner_tree
Returns a tree of the same structure as params where each leaf is the name of the sublearner to apply.
axlearn/common/learner.py:434
↓ 4 callersFunction_locate_user_config_file
Looks for the user's config file in the search paths, or returns None if not found. A user config file may not exist if e.g. the user has never i
axlearn/cloud/common/config.py:162
↓ 4 callersFunction_log_per_layer_stats
Expand the Nested Tensor `stats` and add summaries. Args: stats: A Nested Tensor, e.g., containing param norms or gradient statistics.
axlearn/common/optimizers.py:398
↓ 4 callersFunction_make_bf16_storage
( *, batch: int = 2, num_heads: int = 3, num_pages: int = 8, page_size: int = 4, head_
axlearn/common/kv_cache/paged_kv_storage_test.py:21
↓ 4 callersFunction_make_index_map
Creates an index map function for query/bias tensor.
axlearn/common/flash_attention/tpu_paged_attention_kernel.py:138
↓ 4 callersMethod_make_job
(self, *, enable_replica_restart: Optional[bool] = None, max_tries: int = 10)
axlearn/cloud/gcp/job_test.py:166
↓ 4 callersMethod_maybe_add_volume_mount
(self, volume_mounts: list[dict], *, spec: Optional[VolumeMount])
axlearn/cloud/gcp/jobset_utils.py:614
↓ 4 callersFunction_moment
( val: Tensor, norm_ema: Tensor, norm_square_ema: Tensor, coun
axlearn/common/optimizers.py:1346
↓ 4 callersFunction_mult_to_arg
(level: float, multiplier: float = 1.0)
axlearn/vision/augment.py:644
↓ 4 callersFunction_node_pool_resource
Builds gcloud v1 node pool API resource. Args: credentials: Gcloud credentials used by googleapiclient.discovery. Returns: d
axlearn/cloud/gcp/node_pool.py:234
↓ 4 callersMethod_nonzero
Returns an sequence of biases in this collection except those detected as zero. Returned biases are not guaranteed to be nonzero, but are gua
axlearn/common/attention_bias.py:253
↓ 4 callersFunction_num_replicas_per_shard
Gets the global replication count for each unique shard.
axlearn/common/array_serialization.py:178
↓ 4 callersFunction_opt_params_from_model_params
(model_params, model_param_specs)
axlearn/common/gradient_monitor_test.py:160
↓ 4 callersFunction_pad_logical_to_physical
Pad logical dataset in preparation for batching. Args: dataset: The dataset to pad. global_batch_size: The size of the global phy
axlearn/common/input_tf_data.py:745
↓ 4 callersFunction_parameters_from_t5x_layer_norm
Imports parameters from T5X Layer Norm param dict into AXLearn RMSNorm model state. Corresponding T5X module is layers.LayerNorm Args:
axlearn/common/t5x_param_converter.py:366
↓ 4 callersFunction_parameters_from_t5x_linear_like
Imports parameters from T5X linear param dict into AXLearn Embedding or Linear model state. Corresponding T5X modules are layers.Embed and layers
axlearn/common/t5x_param_converter.py:391
↓ 4 callersFunction_parameters_from_t5x_transformer_layer
Imports parameters from T5X Transformer param dict into AXLearn TransformerLayer model state. Corresponding T5X layer is network.EncoderLayer and
axlearn/common/t5x_param_converter.py:190
↓ 4 callersFunction_parse_axes
Parses an axis pattern string into a hierarchical tuple structure. Converts patterns like "b t (g k) h" into ('b', 't', ('g', 'k'), 'h'). Ru
axlearn/common/ein_ops.py:322
↓ 4 callersMethod_pjit
Compiles `fn` to run on the device mesh. _pjit can be used for functions with a commonly used signature and partitioning scheme. Subc
axlearn/common/evaler.py:167
↓ 4 callersMethod_pop_element
Pops element from self._current_example_list, returns None if the list is empty.
axlearn/common/input_grain_lm.py:90
↓ 4 callersFunction_prompt
Prompts a user for a value. Args: field: The field name to display. required: If True, prompts until user provides a truthy value
axlearn/cloud/common/config.py:177
↓ 4 callersFunction_repo_root_or_cwd
Gets the repo root from CWD, or return CWD if not in a repo. Note that this repo can be arbitrary, and may not contain `ROOT_MODULE_NAME`.
axlearn/cloud/common/config.py:115
↓ 4 callersMethod_reschedule
Reschedules the jobset onto the appropriate tier. If we can identify that the node pool has incorrect selectors for the current scheduling
axlearn/cloud/gcp/runners/gke.py:697
↓ 4 callersMethod_scale_qk
( self, *, q_proj: Tensor, k_proj: Tensor, query_positions: Tensor,
axlearn/common/attention.py:2100
↓ 4 callersFunction_segment_relative_positions
Computes segment-relative positions from segment_ids. For each segment, positions start from 0. Padding positions (segment_ids == 0) get position
axlearn/audio/encoder_asr.py:299
↓ 4 callersMethod_service_config
( self, *, command: str, **kwargs, )
axlearn/cloud/gcp/k8s_service_test.py:19
↓ 4 callersMethod_source
(self, texts: list[str], max_len: int, batch_size: int = 1)
axlearn/common/input_grain_lm_test.py:694
↓ 4 callersMethod_start_gc_thread
Starts garbage collection (if not already started) in a separate thread.
axlearn/common/checkpointer.py:1063
↓ 4 callersMethod_test_extend_step
Shared extend_step test: compares autoregressive decoding against prefill.
axlearn/common/flash_attention/layer_test.py:842
↓ 4 callersMethod_test_forward_vs_extend_step
Tests that {init,prefill}_states + extend_step is equivalent to forward for `cfg`.
axlearn/common/attention_test.py:4137
↓ 4 callersFunction_thick_left_edge
(coordinate, bound)
axlearn/vision/utils_visualization.py:39
↓ 4 callersFunction_thick_right_edge
(coordinate, bound)
axlearn/vision/utils_visualization.py:42
↓ 4 callersFunction_time_call
Times average execution time for fn call over num_iters after warmup.
axlearn/common/flash_attention/tpu_attention_benchmark.py:54
↓ 4 callersMethod_top_k
Selects top-k experts, using configured topk_fn if available.
axlearn/common/mixture_of_experts.py:757
↓ 4 callersMethod_train_step_input_partition_specs
(self)
axlearn/common/trainer.py:399
↓ 4 callersMethod_update_jobs
Handles state transitions for all jobs. The scheduler is used to determine which jobs to resume/pre-empt based on job priority. The a
axlearn/cloud/common/bastion.py:1378
↓ 4 callersFunction_vertexai_experiment_name_from_output_dir
Creates Vertex AI experiment name from output_dir.
axlearn/cloud/gcp/vertexai_tensorboard.py:29
↓ 4 callersFunction_vocab_cfg
()
axlearn/common/input_glue_test.py:31
↓ 4 callersFunction_wrap_exception
Surfaces any `source_exc` as `target_exc` instead. Users can use multiple contexts to wrapping different `target_exc`, e.g: ``` with (
axlearn/common/file_system.py:24
↓ 4 callersMethod_write_per_step
(self, writer: WandBWriter, step: int)
axlearn/common/summary_writer_test.py:208
↓ 4 callersFunctionadamw_decoupled_learner_config
Build learner using the AdamW optimizer and a cosine lr schedule with linear warmup.
axlearn/experiments/text/gpt/common.py:415
↓ 4 callersFunctionand_masks
Returns a MaskFn that's the intersection of provided MaskFn's.
axlearn/common/attention_bias.py:758
↓ 4 callersMethodapply
(self, prng_key: Tensor, params: NestedTensor)
axlearn/common/layers.py:1443
↓ 4 callersFunctionasync_save_tf_savables
Asynchronously saves TF savables from `value_map` into `dir`. When this call returns, `value_map` can be safely mutated, but saving to `dir` will
axlearn/common/checkpointer.py:203
↓ 4 callersFunctionblend
Blend image1 and image2 using 'factor'. Factor can be above 0.0. A value of 0.0 means only image1 is used. A value of 1.0 means only image2
axlearn/vision/augment.py:268
↓ 4 callersFunctionbuild_cfg
Build test flat-configs (without HF reference).
axlearn/common/deberta_test.py:120
↓ 4 callersFunctionbuild_mask
Builds the block map where True means the block is not fully masked. Args: mask_fn: The attention mask function. q_seq_len: Query
axlearn/common/flash_attention/common.py:34
↓ 4 callersMethodbundle
(self, name: str)
axlearn/cloud/gcp/jobs/launch_test.py:136
↓ 4 callersFunctionbundler_flags
Common bundler flags. Keyword args will be forwarded to flag definitions.
axlearn/cloud/common/bundler.py:696
↓ 4 callersMethodcanonicalize
Returns a FanAxes equivalent to this one where all fields are tuples.
axlearn/common/param_init.py:47
↓ 4 callersFunctioncast_floats
Maps valid float arrays found in the inputs to the requested dtype in {float32, bfloat16}. Args: in_tree: The input values. to_dt
axlearn/common/utils.py:1139
↓ 4 callersFunctioncheck_numerics
Checks that all elements in `x` are finite.
axlearn/common/utils.py:359
↓ 4 callersFunctionclip_by_block_rms
Clip updates to a max rms for the gradient of each param vector or matrix. A `block` is here a weight vector (e.g. in a Linear layer) or a weight
axlearn/common/optimizers.py:1509
↓ 4 callersFunctionclip_or_pad_to_fixed_size
Pads data to a fixed length at the first dimension. Args: input_tensor: `Tensor` with any dimension. size: specifies the first di
axlearn/vision/utils_detection.py:929
↓ 4 callersMethodcollect_metrics
Collect metrics from the device, it should be empty.
axlearn/common/monitoring/device_monitor.py:30
↓ 4 callersMethodcompute
Computes a merge matrix by comparing prefixes.
axlearn/audio/decoder_asr.py:119
↓ 4 callersMethodcompute_attention_logit_biases
Produces self-attention logit biases. Args: input_ids: A Tensor of shape [batch_size, seq_len]. segment_ids: An optio
axlearn/common/encoder.py:103
↓ 4 callersFunctionconfig_fn
()
axlearn/experiments/text/gpt/common.py:724
← previousnext →601–700 of 7,799, ranked by callers