mirror of
https://github.com/wassname/ray.git
synced 2026-07-28 11:25:04 +08:00
[tune] Dynamic Resources for Trials (#3974)
## What do these changes do?
Provides a small helper function for modifying the resource requirements of a trial.
Also implements the following:
- setting the last_result to be {} instead of None
- Adding a shuffle to the BasicVariantGenerator
This commit is contained in:
@@ -193,7 +193,7 @@ class Checkpoint(object):
|
||||
def __init__(self, storage, value, last_result=None):
|
||||
self.storage = storage
|
||||
self.value = value
|
||||
self.last_result = last_result
|
||||
self.last_result = last_result or {}
|
||||
|
||||
@staticmethod
|
||||
def from_object(value=None):
|
||||
@@ -283,7 +283,7 @@ class Trial(object):
|
||||
self.max_failures = max_failures
|
||||
|
||||
# Local trial state that is updated during the run
|
||||
self.last_result = None
|
||||
self.last_result = {}
|
||||
self.last_update_time = -float("inf")
|
||||
self.checkpoint_freq = checkpoint_freq
|
||||
self.checkpoint_at_end = checkpoint_at_end
|
||||
@@ -336,6 +336,18 @@ class Trial(object):
|
||||
loggers=self.loggers,
|
||||
sync_function=self.sync_function)
|
||||
|
||||
def update_resources(self, cpu, gpu, **kwargs):
|
||||
"""EXPERIMENTAL: Updates the resource requirements.
|
||||
|
||||
Should only be called when the trial is not running.
|
||||
|
||||
Raises:
|
||||
ValueError if trial status is running.
|
||||
"""
|
||||
if self.status is Trial.RUNNING:
|
||||
raise ValueError("Cannot update resources while Trial is running.")
|
||||
self.resources = Resources(cpu, gpu, **kwargs)
|
||||
|
||||
def sync_logger_to_new_location(self, worker_ip):
|
||||
"""Updates the logger location.
|
||||
|
||||
@@ -392,7 +404,7 @@ class Trial(object):
|
||||
def progress_string(self):
|
||||
"""Returns a progress message for printing out to the console."""
|
||||
|
||||
if self.last_result is None:
|
||||
if not self.last_result:
|
||||
return self._status_string()
|
||||
|
||||
def location_string(hostname, pid):
|
||||
@@ -402,12 +414,12 @@ class Trial(object):
|
||||
return '{} pid={}'.format(hostname, pid)
|
||||
|
||||
pieces = [
|
||||
'{} [{}]'.format(
|
||||
self._status_string(),
|
||||
location_string(
|
||||
self.last_result.get(HOSTNAME),
|
||||
self.last_result.get(PID))), '{} s'.format(
|
||||
int(self.last_result.get(TIME_TOTAL_S)))
|
||||
'{}'.format(self._status_string()), '[{}]'.format(
|
||||
self.resources.summary_string()), '[{}]'.format(
|
||||
location_string(
|
||||
self.last_result.get(HOSTNAME),
|
||||
self.last_result.get(PID))), '{} s'.format(
|
||||
int(self.last_result.get(TIME_TOTAL_S)))
|
||||
]
|
||||
|
||||
if self.last_result.get(TRAINING_ITERATION) is not None:
|
||||
|
||||
Reference in New Issue
Block a user