Cifar dataloader does not work
まだ誰も着手していません。
評価
- 難易度
- 4/5
- 見積もり時間
- 3〜5日
- 初心者へのやさしさ
- 35/100
- issue の種類
- バグ
- 明瞭さ
- おおむね明確
- 活発さ
- 停滞
- 技術スタック
- python
- 領域
- data, machine-learning
調査の方向性
algorithms/baselines/self_tuning/jax_nadamw_full_budget.py から始め、特に get_batch_size と data_selection を確認してから、submission_runner.py コマンドと上記で説明されている CIFAR ワークロードを使って失敗を再現します。batch が Flax prefetch_to_device に到達するまでの流れを追跡し、CIFAR の JAX 実行が shard と device の長さの不一致なしに完了することを確認します。
索引モデルが issue の本文から書いたものです。
説明
The cifar dataloader no longer works properly with jax algorithms using jax.jit. I did not test to see if pytorch algorithms still work with cifar.
Description
When running jax_nadamw_full_budget.py optimizer with the cifar workload, an error is thrown which says
len(shards) = 128 but len(devices) = 8
Here is the relevant log:
I0918 18:15:28.313870 140360727609984 submission_runner.py:359] Starting training loop.
Traceback (most recent call last):
File "/algorithmic-efficiency/submission_runner.py", line 869, in
app.run(main)
File "/usr/local/lib/python3.11/site-packages/absl/app.py", line 308, in run
_run_main(main, args)
File "/usr/local/lib/python3.11/site-packages/absl/app.py", line 254, in _run_main
sys.exit(main(argv))
^^^^^^^^^^
File "/algorithmic-efficiency/submission_runner.py", line 834, in main
score = score_submission_on_workload(
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/algorithmic-efficiency/submission_runner.py", line 747, in score_submission_on_workload
score, _ = train_once(
^^^^^^^^^^^
File "/algorithmic-efficiency/submission_runner.py", line 375, in train_once
batch = data_selection(
^^^^^^^^^^^^^^^
File "/algorithmic-efficiency/algorithms/baselines/self_tuning/jax_nadamw_full_budget.py", line 446, in data_selection
batch = next(input_queue)
^^^^^^^^^^^^^^^^^
File "/usr/local/lib/python3.11/site-packages/flax/jax_utils.py", line 147, in prefetch_to_device
enqueue(size) # Fill up the buffer.
^^^^^^^^^^^^^
File "/usr/local/lib/python3.11/site-packages/flax/jax_utils.py", line 145, in enqueue
queue.append(jax.tree_util.tree_map(_prefetch, data))
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/usr/local/lib/python3.11/site-packages/jax/_src/tree_util.py", line 361, in tree_map
return treedef.unflatten(f(*xs) for xs in zip(*all_leaves))
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/usr/local/lib/python3.11/site-packages/jax/_src/tree_util.py", line 361, in
return treedef.unflatten(f(*xs) for xs in zip(*all_leaves))
^^^^^^
File "/usr/local/lib/python3.11/site-packages/flax/jax_utils.py", line 141, in _prefetch
return jax.device_put_sharded(list(xs), devices)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/usr/local/lib/python3.11/site-packages/jax/_src/api.py", line 2636, in device_put_sharded
raise ValueError(f"len(shards) = {len(shards)} must equal "
ValueError: len(shards) = 128 must equal len(devices) = 8.
2025-09-18 18:15:29.477815: W tensorflow/core/kernels/data/cache_dataset_ops.cc:916] The calling iterator did not fully read the dataset being cached. In order to avoid unexpected truncation of the dataset, the partially cached contents of the dataset will be discarded. This can happen if you have an input pipeline similar todataset.cache().take(k).repeat(). You should usedataset.take(k).cache().repeat()instead.
Steps to Reproduce
- In algorithms/baselines/self_tuning/jax_nadamw_full_budget.py add the following two lines in get_batch_size function:
elif workload_name == 'cifar':
return 128
- Then run the cifar workload in docker:
python submission_runner.py
--framework=jax
--workload=cifar
--experiment_dir=/experiment_runs
--experiment_name=jax_debug_cifar
--data_dir=/data
--tuning_ruleset=self
--submission_path=algorithms/baselines/self_tuning/jax_nadamw_full_budget.py
Source or Possible Fix
I think the cifar is not an officially supported workload, but it can be useful for debugging. So once it is not too much trouble we should fix this.
- 主要言語
- Python
- スター
- 425
- フォーク
- 79
- PR マージ指標
- 30日以内にマージされた PR はありません
環境構築
- Dockerfile・Docker Compose ファイルなし
- プルリクエストのテンプレートなし
- コントリビューションガイドを読む
はじめの一歩
- issue を最後まで読み、次にプロジェクトのコントリビューションガイドを読みます。
- 着手することを issue にコメントします — 二人が同じ作業をするのを防げます。
- リポジトリをフォークし、ブランチを切って変更します。
- issue 番号を参照したプルリクエストを送ります。
mlcommons/algorithmic-efficiency のほかの issue
-
難易度 5/5 1週間以上 初心者へのやさしさ 30/100
-
難易度 4/5 3〜5日 初心者へのやさしさ 38/100
mlcommons/algorithmic-efficiency#924 · コメント 5 件 ·
-
Anima: live PH monitoring in conversational agent — overfitting detection every 50 interactionsオープン
難易度 5/5 1週間以上 初心者へのやさしさ 25/100
-
難易度 5/5 1週間以上 初心者へのやさしさ 20/100
-
Topological overfitting detection: H0 gap catches overfitting before accuracy diverges (r=0.998)オープン
難易度 5/5 1週間以上 初心者へのやさしさ 15/100
mlcommons/algorithmic-efficiency の issue をすべて見る
似ている issue
-
難易度 2/5 1〜3時間 初心者へのやさしさ 78/100
conda-forge/conda-build-feedstock#289 · コメント 1 件 · リアクション 1 件 ·
-
`pulptest` no longer works in 4.0.0: `ImportError: Start directory is not importable: 'pulp/tests'`オープン
難易度 2/5 1〜3時間 初心者へのやさしさ 75/100
-
難易度 2/5 1〜3時間 初心者へのやさしさ 72/100
-
難易度 2/5 1〜3時間 初心者へのやさしさ 78/100
-
remove reddit feedsオープン
難易度 2/5 1〜3時間 初心者へのやさしさ 72/100
TomCasavant/ohio-sites#224 ·