[BUG] GkeTpuCluster queries metadata.google.internal for MEGASCALE_SLICE_ID even when TPU_SKIP_MDS_QUERY is set

Ouverte Adaptée aux débutants
#40,854 0 commentaires 0 réactions 0 personnes assignées Voir sur GitHub

Personne n'a encore pris cette issue.

Évaluation

Difficulté
2/5
Temps estimé
1-3 heures
Accessibilité débutants
88/100
Type d'issue
Bug
Clarté
Clairement spécifiée
Activité
Active
Stack technique
python

Piste de recherche

Commencez dans jax/_src/clusters/cloud_tpu_cluster.py en lisant get_tpu_env_value() et ses appelants, en particulier GkeTpuCluster.get_process_id(). Reproduisez le problème avec TPU_SKIP_MDS_QUERY=true et sans accès aux métadonnées ; c’est terminé lorsque le serveur de métadonnées n’est pas interrogé et que les valeurs de secours existantes permettent à l’initialisation de se poursuivre.

Rédigé par le modèle d'indexation à partir du texte de l'issue.

Description

Description
Summary

In jax/_src/clusters/cloud_tpu_cluster.py, TPU_SKIP_MDS_QUERY is checked in GceTpuCluster.is_env_present() (line 186) to bypass querying the GCE Metadata Server. However, get_tpu_env_value() (line 70) completely ignores TPU_SKIP_MDS_QUERY:

def get_tpu_env_value(key) -> str | None:
  # First try to get the value from the environment.
  value = os.environ.get(key, None)
  if value is None:
    # If not found, try to get it from the metadata.
    value = get_tpu_env_value_from_metadata(key)
  return value

When running in Kubernetes with GkeTpuCluster, calling GkeTpuCluster.get_process_id() queries MEGASCALE_SLICE_ID via get_tpu_env_value('MEGASCALE_SLICE_ID').

Even when TPU_SKIP_MDS_QUERY=true is explicitly set in the container (e.g. by Kubernetes TPU device plugins, or on bare-metal/testbed clusters), get_tpu_env_value() still unconditionally calls get_tpu_env_value_from_metadata(), which tries to reach http://metadata.google.internal/computeMetadata/v1/instance/attributes/tpu-env.

When the metadata server is unreachable or unresolvable, JAX hangs for 6 retries with 60-second timeouts and then crashes with requests.exceptions.ConnectionError.

Notice that the caller functions in BaseTpuCluster already have built-in defaults if get_tpu_env_value() returns None:

  • _get_slice_id(): if not slice_id: return 0
  • _get_num_slices(): if not num_slices: return 1
  • get_coordinator_address(): if not coordinator_address: coordinator_address = cls._get_worker_list_in_slice()[0]

Therefore, simply checking TPU_SKIP_MDS_QUERY in get_tpu_env_value() allows all existing fallbacks to work cleanly without contacting the metadata server.

Steps to Reproduce

Run JAX in a Kubernetes TPU container where TPU_SKIP_MDS_QUERY=true and TPU_WORKER_HOSTNAMES are set, but without access to metadata.google.internal:

import jax
# GkeTpuCluster is selected, but crashes trying to reach metadata.google.internal
jax.distributed.initialize()
Stack Trace
Traceback (most recent call last):
  File "jax/_src/clusters/cloud_tpu_cluster.py", line 146, in get_process_id
    slice_id = cls._get_slice_id()
  File "jax/_src/clusters/cloud_tpu_cluster.py", line 162, in _get_slice_id
    slice_id = get_tpu_env_value('MEGASCALE_SLICE_ID')
  File "jax/_src/clusters/cloud_tpu_cluster.py", line 75, in get_tpu_env_value
    value = get_tpu_env_value_from_metadata(key)
  File "jax/_src/clusters/cloud_tpu_cluster.py", line 45, in get_metadata
    api_resp = requests.get(
requests.exceptions.ConnectionError: HTTPConnectionPool(host='metadata.google.internal', port=80): Max retries exceeded with url: /computeMetadata/v1/instance/attributes/tpu-env (Caused by NameResolutionError("HTTPConnection(host='metadata.google.internal', port=80): Failed to resolve 'metadata.google.internal'"))
Proposed Fix

In jax/_src/clusters/cloud_tpu_cluster.py:

def get_tpu_env_value(key) -> str | None:
  # First try to get the value from the environment.
  value = os.environ.get(key, None)
  if value is None and os.environ.get("TPU_SKIP_MDS_QUERY") is None:
    # If not found, try to get it from the metadata.
    value = get_tpu_env_value_from_metadata(key)
  return value
System info (python version, jaxlib version, accelerator, etc.)
  • OS: Linux
  • Accelerator: Google Cloud TPU
  • JAX version: main / 0.4.x+
  • Environment: Kubernetes / Bare-metal TPU cluster with TPU_SKIP_MDS_QUERY=true
Langage dominant
Python
Étoiles
36.3k
Forks
3.8k
Merge moyen
1 j 2 h
PR mergées (30 j)
402

Guide de contribution

Ouvrir le guide de contribution

Par où commencer

  1. Lisez l'issue en entier, puis le guide de contribution du projet.
  2. Signalez en commentaire que vous la prenez — cela évite que deux personnes fassent le même travail.
  3. Forkez le dépôt et travaillez sur une branche.
  4. Ouvrez une pull request qui référence le numéro de l'issue.

Autres issues de jax-ml/jax

Toutes les issues de jax-ml/jax

Issues similaires

Plus d'issues Python

Recevez les nouvelles issues par e-mail

Un résumé court des issues GitHub adaptées aux débutants.