Lightning-AI / Lightning-AI/pytorch-lightning

Training crash when using XLA profiler on XLA accelerator and manual optimization

Open
#20,206 0 comments 0 reactions 0 assignees View on GitHub

Nobody has claimed this yet.

bug strategy: xla ver: 2.4.x
Dominant language
Python
Stars
31.4k
Forks
3.8k
Avg merge
6d 7h
Merged PRs (30d)
6

Description

### Bug description

training loop crash when running on XLA profiler + manual optimization.

### What version are you seeing the problem on?

v2.4

### How to reproduce the bug

```python
Training on XLAProfile + Manual Optimization on XLA Machine
```

### Error messages and logs

```

concurrent.futures.process._RemoteTraceback:
"""
Traceback (most recent call last):
File "/~/miniconda3/envs/ldm-tp23/lib/python3.10/concurrent/futures/process.py", line 246, in _process_worker
r = call_item.fn(*call_item.args, **call_item.kwargs)
File "/~/miniconda3/envs/ldm-tp23/lib/python3.10/concurrent/futures/process.py", line 205, in _process_chunk
return [fn(*args) for args in chunk]
File "/~/miniconda3/envs/ldm-tp23/lib/python3.10/concurrent/futures/process.py", line 205, in
return [fn(*args) for args in chunk]
File "/~/miniconda3/envs/ldm-tp23/lib/python3.10/site-packages/torch_xla/runtime.py", line 95, in wrapper
return fn(*args, **kwargs)
File "/~/miniconda3/envs/ldm-tp23/lib/python3.10/site-packages/torch_xla/_internal/pjrt.py", line 78, in _run_thread_per_device
replica_results = list(
File "/~/miniconda3/envs/ldm-tp23/lib/python3.10/concurrent/futures/_base.py", line 621, in result_iterator
yield _result_or_cancel(fs.pop())
File "/~/miniconda3/envs/ldm-tp23/lib/python3.10/concurrent/futures/_base.py", line 319, in _result_or_cancel
return fut.result(timeout)
File "/~/miniconda3/envs/ldm-tp23/lib/python3.10/concurrent/futures/_base.py", line 458, in result
return self.__get_result()
File "/~/miniconda3/envs/ldm-tp23/lib/python3.10/concurrent/futures/_base.py", line 403, in __get_result
raise self._exception
File "/~/miniconda3/envs/ldm-tp23/lib/python3.10/concurrent/futures/thread.py", line 58, in run
result = self.fn(*self.args, **self.kwargs)
File "/~/miniconda3/envs/ldm-tp23/lib/python3.10/site-packages/torch_xla/_internal/pjrt.py", line 71, in _thread_fn
return fn()
File "/~/miniconda3/envs/ldm-tp23/lib/python3.10/site-packages/torch_xla/_internal/pjrt.py", line 187, in __call__
self.fn(runtime.global_ordinal(), *self.args, **self.kwargs)
File "/~/miniconda3/envs/ldm-tp23/lib/python3.10/site-packages/pytorch_lightning/strategies/launchers/xla.py", line 141, in _wrapping_function
results = function(*args, **kwargs)
File "/~/miniconda3/envs/ldm-tp23/lib/python3.10/site-packages/pytorch_lightning/trainer/trainer.py", line 579, in _fit_impl
self._run(model, ckpt_path=ckpt_path)
File "/~/miniconda3/envs/ldm-tp23/lib/python3.10/site-packages/pytorch_lightning/trainer/trainer.py", line 986, in _run
results = self._run_stage()
File "/~/miniconda3/envs/ldm-tp23/lib/python3.10/site-packages/pytorch_lightning/trainer/trainer.py", line 1030, in _run_stage
self.fit_loop.run()
File "/~/miniconda3/envs/ldm-tp23/lib/python3.10/site-packages/pytorch_lightning/loops/fit_loop.py", line 205, in run
self.advance()
File "/~/miniconda3/envs/ldm-tp23/lib/python3.10/site-packages/pytorch_lightning/loops/fit_loop.py", line 363, in advance
self.epoch_loop.run(self._data_fetcher)
File "/~/miniconda3/envs/ldm-tp23/lib/python3.10/site-packages/pytorch_lightning/loops/training_epoch_loop.py", line 140, in run
self.advance(data_fetcher)
File "/~/miniconda3/envs/ldm-tp23/lib/python3.10/site-packages/pytorch_lightning/loops/training_epoch_loop.py", line 252, in advance
batch_output = self.manual_optimization.run(kwargs)
File "/~/miniconda3/envs/ldm-tp23/lib/python3.10/site-packages/pytorch_lightning/loops/optimization/manual.py", line 94, in run
self.advance(kwargs)
File "/~/miniconda3/envs/ldm-tp23/lib/python3.10/site-packages/pytorch_lightning/loops/optimization/manual.py", line 114, in advance
training_step_output = call._call_strategy_hook(trainer, "training_step", *kwargs.values())
File "/~/miniconda3/envs/ldm-tp23/lib/python3.10/site-packages/pytorch_lightning/trainer/call.py", line 311, in _call_strategy_hook
output = fn(*args, **kwargs)
File "/~/miniconda3/envs/ldm-tp23/lib/python3.10/site-packages/pytorch_lightning/strategies/strategy.py", line 390, in training_step
return self.lightning_module.training_step(*args, **kwargs)
File "/mnt/disks/persist/ldm/ldm/models/autoencoder.py", line 438, in training_step
opt1.step()
File "/~/miniconda3/envs/ldm-tp23/lib/python3.10/site-packages/pytorch_lightning/core/optimizer.py", line 153, in step
step_output = self._strategy.optimizer_step(self._optimizer, closure, **kwargs)
File "/~/miniconda3/envs/ldm-tp23/lib/python3.10/site-packages/pytorch_lightning/strategies/ddp.py", line 270, in optimizer_step
optimizer_output = super().optimizer_step(optimizer, closure, model, **kwargs)
File "/~/miniconda3/envs/ldm-tp23/lib/python3.10/site-packages/pytorch_lightning/strategies/strategy.py", line 238, in optimizer_step
return self.precision_plugin.optimizer_step(optimizer, model=model, closure=closure, **kwargs)
File "/~/miniconda3/envs/ldm-tp23/lib/python3.10/site-packages/pytorch_lightning/plugins/precision/xla.py", line 75, in optimizer_step
xm.mark_step()
File "/~/miniconda3/envs/ldm-tp23/lib/python3.10/site-packages/torch_xla/core/xla_model.py", line 1056, in mark_step
torch_xla._XLAC._xla_step_marker(
RuntimeError: Expecting scope to be empty but it is [Strategy]XLAStrategy.training_step.1
Exception raised from ResetScopeContext at ../torch/csrc/lazy/core/ir_metadata.cpp:77 (most recent call first):
frame #0: c10::Error::Error(c10::SourceLocation, std::string) + 0x57 (0x7f812737a897 in /~/miniconda3/envs/ldm-tp23/lib/python3.10/site-packages/torch/lib/libc10.so)
frame #1: c10::detail::torchCheckFail(char const*, char const*, unsigned int, std::string const&) + 0x64 (0x7f812732ab25 in /~/miniconda3/envs/ldm-tp23/lib/python3.10/site-packages/torch/lib/libc10.so)
frame #2: torch::lazy::ScopePusher::ResetScopes() + 0xa5 (0x7f81136f7c55 in /~/miniconda3/envs/ldm-tp23/lib/python3.10/site-packages/torch/lib/libtorch_cpu.so)
frame #3: torch_xla::XLAGraphExecutor::MarkStep(torch::lazy::BackendDevice const&) + 0x57 (0x7f7fc6920a87 in /~/miniconda3/envs/ldm-tp23/lib/python3.10/site-packages/_XLAC.cpython-310-x86_64-linux-gnu.so)
frame #4: + 0x4aeb60a (0x7f7fc66eb60a in /~/miniconda3/envs/ldm-tp23/lib/python3.10/site-packages/_XLAC.cpython-310-x86_64-linux-gnu.so)
frame #5: + 0x4aebab6 (0x7f7fc66ebab6 in /~/miniconda3/envs/ldm-tp23/lib/python3.10/site-packages/_XLAC.cpython-310-x86_64-linux-gnu.so)
frame #6: + 0x4abd006 (0x7f7fc66bd006 in /~/miniconda3/envs/ldm-tp23/lib/python3.10/site-packages/_XLAC.cpython-310-x86_64-linux-gnu.so)
frame #7: python() [0x4fdc87]

frame #12: python() [0x5099ce]
frame #15: python() [0x509b26]
frame #17: python() [0x509b26]
frame #19: python() [0x5099ce]
frame #21: python() [0x509b26]
frame #23: python() [0x509b26]
frame #41: python() [0x5099ce]
frame #43: python() [0x509b26]
frame #45: python() [0x509b26]
frame #49: python() [0x5cf883]
frame #51: python() [0x5c87f7]

"""

The above exception was the direct cause of the following exception:

Traceback (most recent call last):
File "/mnt/disks/persist/ldm/main.py", line 753, in
trainer.fit(model, data, ckpt_path=opt.resume_from_checkpoint if "resume_from_checkpoint" in opt else None)
File "/~/miniconda3/envs/ldm-tp23/lib/python3.10/site-packages/pytorch_lightning/trainer/trainer.py", line 543, in fit
call._call_and_handle_interrupt(
File "/~/miniconda3/envs/ldm-tp23/lib/python3.10/site-packages/pytorch_lightning/trainer/call.py", line 43, in _call_and_handle_interrupt
return trainer.strategy.launcher.launch(trainer_fn, *args, trainer=trainer, **kwargs)
File "/~/miniconda3/envs/ldm-tp23/lib/python3.10/site-packages/pytorch_lightning/strategies/launchers/xla.py", line 98, in launch
process_context = xmp.spawn(
File "/~/miniconda3/envs/ldm-tp23/lib/python3.10/site-packages/torch_xla/runtime.py", line 95, in wrapper
return fn(*args, **kwargs)
File "/~/miniconda3/envs/ldm-tp23/lib/python3.10/site-packages/torch_xla/distributed/xla_multiprocessing.py", line 38, in spawn
return pjrt.spawn(fn, nprocs, start_method, args)
File "/~/miniconda3/envs/ldm-tp23/lib/python3.10/site-packages/torch_xla/_internal/pjrt.py", line 211, in spawn
run_multiprocess(spawn_fn, start_method=start_method)
File "/~/miniconda3/envs/ldm-tp23/lib/python3.10/site-packages/torch_xla/runtime.py", line 95, in wrapper
return fn(*args, **kwargs)
File "/~/miniconda3/envs/ldm-tp23/lib/python3.10/site-packages/torch_xla/_internal/pjrt.py", line 171, in run_multiprocess
replica_results = list(
File "/~/miniconda3/envs/ldm-tp23/lib/python3.10/site-packages/torch_xla/_internal/pjrt.py", line 172, in
itertools.chain.from_iterable(
File "/~/miniconda3/envs/ldm-tp23/lib/python3.10/concurrent/futures/process.py", line 575, in _chain_from_iterable_of_lists
for element in iterable:
File "/~/miniconda3/envs/ldm-tp23/lib/python3.10/concurrent/futures/_base.py", line 621, in result_iterator
yield _result_or_cancel(fs.pop())
File "/~/miniconda3/envs/ldm-tp23/lib/python3.10/concurrent/futures/_base.py", line 319, in _result_or_cancel
return fut.result(timeout)
File "/~/miniconda3/envs/ldm-tp23/lib/python3.10/concurrent/futures/_base.py", line 458, in result
return self.__get_result()
File "/~/miniconda3/envs/ldm-tp23/lib/python3.10/concurrent/futures/_base.py", line 403, in __get_result
raise self._exception
RuntimeError: Expecting scope to be empty but it is [Strategy]XLAStrategy.training_step.1
Exception raised from ResetScopeContext at ../torch/csrc/lazy/core/ir_metadata.cpp:77 (most recent call first):
frame #0: c10::Error::Error(c10::SourceLocation, std::string) + 0x57 (0x7f812737a897 in /~/miniconda3/envs/ldm-tp23/lib/python3.10/site-packages/torch/lib/libc10.so)
frame #1: c10::detail::torchCheckFail(char const*, char const*, unsigned int, std::string const&) + 0x64 (0x7f812732ab25 in /~/miniconda3/envs/ldm-tp23/lib/python3.10/site-packages/torch/lib/libc10.so)
frame #2: torch::lazy::ScopePusher::ResetScopes() + 0xa5 (0x7f81136f7c55 in /~/miniconda3/envs/ldm-tp23/lib/python3.10/site-packages/torch/lib/libtorch_cpu.so)
frame #3: torch_xla::XLAGraphExecutor::MarkStep(torch::lazy::BackendDevice const&) + 0x57 (0x7f7fc6920a87 in /~/miniconda3/envs/ldm-tp23/lib/python3.10/site-packages/_XLAC.cpython-310-x86_64-linux-gnu.so)
frame #4: + 0x4aeb60a (0x7f7fc66eb60a in /~/miniconda3/envs/ldm-tp23/lib/python3.10/site-packages/_XLAC.cpython-310-x86_64-linux-gnu.so)
frame #5: + 0x4aebab6 (0x7f7fc66ebab6 in /~/miniconda3/envs/ldm-tp23/lib/python3.10/site-packages/_XLAC.cpython-310-x86_64-linux-gnu.so)
frame #6: + 0x4abd006 (0x7f7fc66bd006 in /~/miniconda3/envs/ldm-tp23/lib/python3.10/site-packages/_XLAC.cpython-310-x86_64-linux-gnu.so)
frame #7: python() [0x4fdc87]

frame #12: python() [0x5099ce]
frame #15: python() [0x509b26]
frame #17: python() [0x509b26]
frame #19: python() [0x5099ce]
frame #21: python() [0x509b26]
frame #23: python() [0x509b26]
frame #41: python() [0x5099ce]
frame #43: python() [0x509b26]
frame #45: python() [0x509b26]
frame #49: python() [0x5cf883]
frame #51: python() [0x5c87f7]
```

### Environment

Current environment

```
- PyTorch Lightning Version (2.4.0):
- PyTorch XLA Version (2.4.0):
- PyTorch Version (2.4):
- Python version (3.10):
```

### More info

_No response_

cc @JackCaoG @Liyang90 @gkroiz

Contributor guide

Open the contributing guide

First steps

  1. Read the whole issue, then the project's contributing guide.
  2. Comment on the issue to say you are picking it up — it saves two people doing the same work.
  3. Fork the repository and make your change on a branch.
  4. Open a pull request that references the issue number.

Research direction

Start by reading the XLA integration points named in the traceback: strategies/launchers/xla.py and plugins/precision/xla.py, alongside loops/optimization/manual.py. Reproduce the reported combination of XLA profiling and manual optimization on the stated versions, then verify that training no longer raises the scope error and add a regression test if the repository has coverage for this path.

Written by the indexing model from the issue text.

Assessment

Tech stack
python, pytorch
Domain
distributed-systems, machine-learning
Issue type
Bug
Difficulty
4/5
Estimated time
3-5 days
Activity status
Stale
Clarity
Needs clarification
Newbie friendliness
25/100

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.