[BUG] GkeTpuCluster queries metadata.google.internal for MEGASCALE_SLICE_ID even when TPU_SKIP_MDS_QUERY is set
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
- Domaine
- cloud, distributed-systems
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 1get_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
- Lisez l'issue en entier, puis le guide de contribution du projet.
- Signalez en commentaire que vous la prenez — cela évite que deux personnes fassent le même travail.
- Forkez le dépôt et travaillez sur une branche.
- Ouvrez une pull request qui référence le numéro de l'issue.
Autres issues de jax-ml/jax
-
Difficulté 1/5 Moins d'une heure Accessibilité débutants 92/100
-
Difficulté 2/5 1-3 heures Accessibilité débutants 76/100
-
bug
Difficulté 2/5 1-3 heures Accessibilité débutants 84/100
-
Difficulté 2/5 1-3 heures Accessibilité débutants 84/100
-
bug
Difficulté 2/5 1-3 heures Accessibilité débutants 70/100
Toutes les issues de jax-ml/jax
Issues similaires
-
Difficulté 2/5 1-3 heures Accessibilité débutants 82/100
-
Difficulté 2/5 1-3 heures Accessibilité débutants 78/100
-
enhancement
Difficulté 2/5 1-3 heures Accessibilité débutants 72/100
-
Difficulté 2/5 1-3 heures Accessibilité débutants 74/100
-
Difficulté 2/5 1-3 heures Accessibilité débutants 84/100
PolicyEngine/policyengine-us#9559 ·