google-deepmind / google-deepmind/deepmind-research
Perceiver IO: "Error while trying to execute the imagenet training script"
- Dominant language
- Jupyter Notebook
- Stars
- 15.2k
- Forks
- 2.9k
- PR merge metrics
- No merged PRs in 30d
Description
Please help, I have been getting the error below while trying to execute
nvidia-smi and nvcc -V return the right values.
************************** Error ************************************
I1202 09:24:15.404018 140515002271552 train.py:70] Training with config:
best_model_eval_metric: eval_top_1_acc
best_model_eval_metric_higher_is_better: true
checkpoint_dir: /tmp/perceiver_imagnet_checkpoints
eval_initial_weights: false
eval_specific_checkpoint_dir: ''
experiment_kwargs:
config:
data:
augmentation:
cutmix: true
mixup_alpha: 0.2
randaugment:
magnitude: 5
num_layers: 4
im_dim: 32
num_classes: 1000
evaluation:
batch_size: 2
subset: test
model:
perceiver_kwargs:
decoder:
num_z_channels: 1024
position_encoding_type: trainable
trainable_position_encoding_kwargs:
init_scale: 0.02
num_channels: 1024
use_query_residual: true
encoder:
cross_attend_widening_factor: 1
cross_attention_shape_for_attn: kv
dropout_prob: 0.0
num_blocks: 2
num_cross_attend_heads: 1
num_self_attend_heads: 8
num_self_attends_per_block: 2
num_z_channels: 1024
self_attend_widening_factor: 1
use_query_residual: true
z_index_dim: 512
z_pos_enc_init_scale: 0.02
input_preprocessor:
concat_or_add_pos: concat
fourier_position_encoding_kwargs:
concat_pos: true
max_resolution: !!python/tuple
- 224
- 224
num_bands: 64
sine_only: false
num_channels: 64
position_encoding_type: fourier
prep_type: pixels
project_pos_dim: -1
spatial_downsample: 1
trainable_position_encoding_kwargs:
init_scale: 0.02
num_channels: 258
optimizer:
adam_kwargs:
b1: 0.9
b2: 0.999
eps: 1.0e-08
base_lr: 0.0005
constant_cosine_decay_kwargs:
constant_fraction: 0.5
end_value: 0.0
cosine_decay_kwargs:
end_value: 0.0
init_value: 0.0
warmup_epochs: 0
decay_pos_embs: true
lamb_kwargs:
b1: 0.9
b2: 0.999
eps: 1.0e-06
max_norm: 10.0
optimizer: lamb
scale_by_batch: true
schedule_type: constant_cosine
step_decay_kwargs:
decay_boundaries:
- 0.5
- 0.8
- 0.95
decay_rate: 0.1
weight_decay: 0.1
training:
batch_size: 2
images_per_epoch: 1281167
label_smoothing: 0.1
n_epochs: 110
interval_type: secs
log_all_train_data: false
log_tensors_interval: 60
log_train_data_interval: 60.0
max_checkpoints_to_keep: 5
n_epochs: 110
one_off_evaluate: false
random_mode_eval: same_host_same_device
random_mode_train: unique_host_unique_device
random_seed: 42
save_checkpoint_interval: 300
train_batch_size: 2
train_checkpoint_all_hosts: false
training_steps: 70464185
/home/jamy/addons/anaconda3/envs/my_env/lib/python3.8/site-packages/jax/_src/lib/xla_bridge.py:399: UserWarning: jax.host_id has been renamed to jax.process_index. This alias will eventually be removed; please update your code.
warnings.warn(
I1202 09:24:15.409622 140515002271552 xla_bridge.py:230] Unable to initialize backend 'tpu_driver': NOT_FOUND: Unable to find driver in registry given worker:
I1202 09:24:15.410672 140515002271552 xla_bridge.py:230] Unable to initialize backend 'tpu': INVALID_ARGUMENT: TpuPlatform is not available.
I1202 09:24:15.579092 140515002271552 utils.py:234] [jaxline] experiment init starting...
I1202 09:24:15.579470 140515002271552 utils.py:241] [jaxline] experiment init finished.
Traceback (most recent call last):
File "perceiver/train/experiment.py", line 538, in
app.run(functools.partial(platform.main, Experiment))
File "/home/jamy/addons/anaconda3/envs/my_env/lib/python3.8/site-packages/absl/app.py", line 312, in run
_run_main(main, args)
File "/home/jamy/addons/anaconda3/envs/my_env/lib/python3.8/site-packages/absl/app.py", line 258, in _run_main
sys.exit(main(argv))
File "/home/jamy/addons/anaconda3/envs/my_env/lib/python3.8/site-packages/jaxline/utils.py", line 401, in inner_wrapper
return f(*args, **kwargs)
File "/home/jamy/addons/anaconda3/envs/my_env/lib/python3.8/site-packages/jaxline/platform.py", line 126, in main
train.train(experiment_class, config, checkpointer, writer)
File "/home/jamy/addons/anaconda3/envs/my_env/lib/python3.8/site-packages/jaxline/utils.py", line 529, in inner_wrapper
return fn(*args, **kwargs)
File "/home/jamy/addons/anaconda3/envs/my_env/lib/python3.8/site-packages/jaxline/train.py", line 81, in train
state.train_step_rng = utils.bcast_local_devices(rng)
File "/home/jamy/addons/anaconda3/envs/my_env/lib/python3.8/site-packages/jaxline/utils.py", line 148, in bcast_local_devices
return jax.tree_map(
File "/home/jamy/addons/anaconda3/envs/my_env/lib/python3.8/site-packages/jax/_src/tree_util.py", line 178, in tree_map
return treedef.unflatten(f(*xs) for xs in zip(*all_leaves))
File "/home/jamy/addons/anaconda3/envs/my_env/lib/python3.8/site-packages/jax/_src/tree_util.py", line 178, in
return treedef.unflatten(f(*xs) for xs in zip(*all_leaves))
File "/home/jamy/addons/anaconda3/envs/my_env/lib/python3.8/site-packages/jaxline/utils.py", line 149, in
lambda v: jax.api.device_put_sharded(len(devices) * [v], devices), value)
AttributeError: module 'jax' has no attribute 'api'
Contributor guide
Assessment
This issue has not been assessed yet.