spec-b300 / source /scripts /cluster /test_ngram_interactive_supervisor.py
khazic's picture
Archive three-epoch run: logs and provenance part 2
932bc69 verified
Raw History Blame Contribute Delete
6.32 kB
"""CPU-only tests: renewal uses salloc; cancellation/failure never renews."""
import shlex
import subprocess
import tempfile
import unittest
from pathlib import Path
from unittest.mock import MagicMock, patch
from ngram_interactive_supervisor import Supervisor
class SupervisorTests(unittest.TestCase):
def controller(self, root):
item = Supervisor.__new__(Supervisor)
item.root = root
item.source = root / "source"
item.job = "1"
item.ssh = ["ssh"]
item.login = "login"
item.status = MagicMock()
item.child = MagicMock()
item.child.poll.return_value = 1
item.child.returncode = 1
item.maintenance = None
item.start_segment = MagicMock(return_value=root / "segments/1")
return item
def test_timeout_renews(self):
with tempfile.TemporaryDirectory() as directory:
controller = self.controller(Path(directory))
controller.job_state = MagicMock(return_value="TIMEOUT")
controller.allocate = MagicMock(side_effect=RuntimeError("test reached renewal"))
with self.assertRaisesRegex(RuntimeError, "reached renewal"):
controller.run()
controller.allocate.assert_called_once()
def test_cancellation_does_not_renew(self):
with tempfile.TemporaryDirectory() as directory:
controller = self.controller(Path(directory))
controller.job_state = MagicMock(return_value="CANCELLED")
controller.allocate = MagicMock()
with self.assertRaisesRegex(RuntimeError, "not renewing"):
controller.run()
controller.allocate.assert_not_called()
def test_long_completing_then_timeout_renews(self):
with tempfile.TemporaryDirectory() as directory:
controller = self.controller(Path(directory))
controller.job_state = MagicMock(
side_effect=["COMPLETING"] * 20 + ["UNKNOWN", "TIMEOUT"]
)
controller.allocate = MagicMock(side_effect=RuntimeError("test reached renewal"))
with patch("ngram_interactive_supervisor.time.sleep") as sleep:
with self.assertRaisesRegex(RuntimeError, "reached renewal"):
controller.run()
self.assertEqual(sleep.call_count, 21)
controller.allocate.assert_called_once()
def test_accounting_timeout_overrides_completing(self):
with tempfile.TemporaryDirectory() as directory:
controller = self.controller(Path(directory))
controller.remote = MagicMock(side_effect=["COMPLETING", "TIMEOUT|"])
self.assertEqual(controller.job_state(), "TIMEOUT")
def test_expired_job_missing_from_squeue_uses_accounting(self):
with tempfile.TemporaryDirectory() as directory:
controller = self.controller(Path(directory))
controller.remote = MagicMock(side_effect=[
subprocess.CalledProcessError(1, "squeue"), "TIMEOUT|"
])
self.assertEqual(controller.job_state(), "TIMEOUT")
def test_transient_query_error_retries(self):
with tempfile.TemporaryDirectory() as directory:
controller = self.controller(Path(directory))
controller.job_state = MagicMock(side_effect=[
subprocess.TimeoutExpired("ssh", 30), "TIMEOUT"
])
with patch("ngram_interactive_supervisor.time.sleep"):
self.assertEqual(controller.wait_for_allocation_exit(), "TIMEOUT")
def test_dead_trainer_in_running_allocation_does_not_renew(self):
with tempfile.TemporaryDirectory() as directory:
controller = self.controller(Path(directory))
controller.job_state = MagicMock(return_value="RUNNING")
controller.allocate = MagicMock()
with patch("ngram_interactive_supervisor.time.monotonic", side_effect=[0, 61]):
self.assertEqual(controller.wait_for_allocation_exit(), "RUNNING")
controller.allocate.assert_not_called()
def test_restart_after_timeout_does_not_relaunch_old_segment(self):
with tempfile.TemporaryDirectory() as directory:
controller = self.controller(Path(directory))
(controller.root / "segments" / controller.job).mkdir(parents=True)
controller.job_state = MagicMock(return_value="TIMEOUT")
controller.allocate = MagicMock(side_effect=RuntimeError("test reached renewal"))
with self.assertRaisesRegex(RuntimeError, "reached renewal"):
controller.run()
controller.start_segment.assert_not_called()
controller.allocate.assert_called_once()
def test_restart_after_cancellation_does_not_renew(self):
with tempfile.TemporaryDirectory() as directory:
controller = self.controller(Path(directory))
(controller.root / "segments" / controller.job).mkdir(parents=True)
controller.job_state = MagicMock(return_value="CANCELLED")
controller.allocate = MagicMock()
with self.assertRaisesRegex(RuntimeError, "not renewing"):
controller.run()
controller.start_segment.assert_not_called()
controller.allocate.assert_not_called()
def test_allocation_command_and_fairshare(self):
with tempfile.TemporaryDirectory() as directory:
controller = self.controller(Path(directory))
(controller.root / "allocations").mkdir()
controller.remote = MagicMock(return_value="low|user|||||0.1|\nbest|user|||||0.9|")
def launch(command, **kwargs):
allocation = shlex.split(command[-1])
self.assertEqual(allocation[0], "salloc")
for flag in ("--account=best", "--nodes=1", "--gpus-per-node=4", "--time=04:00:00", "--qos=interactive"):
self.assertIn(flag, allocation)
(Path(allocation[-1]) / "job_id").write_text("90001\n")
return MagicMock()
with patch("ngram_interactive_supervisor.subprocess.Popen", side_effect=launch):
controller.allocate()
self.assertEqual(controller.job, "90001")
if __name__ == "__main__":
unittest.main()