Lightning-AI / Lightning-AI/pytorch-lightning
Training crash when using XLA profiler on XLA accelerator and manual optimization
Nobody has claimed this yet.
- 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
First steps
- Read the whole issue, then the project's contributing guide.
- Comment on the issue to say you are picking it up — it saves two people doing the same work.
- Fork the repository and make your change on a branch.
- 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