update from last round
This commit is contained in:
+76
-37
@@ -1,6 +1,5 @@
|
||||
from calendar import week
|
||||
import datetime
|
||||
from distutils.log import debug
|
||||
import itertools
|
||||
from typing import Dict, Iterable, List, Sequence, Tuple, Set, Any, Type, Literal
|
||||
import time
|
||||
@@ -9,10 +8,12 @@ from pydantic import BaseModel, Extra, constr
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
import uuid
|
||||
|
||||
import datetime
|
||||
|
||||
from pyomo.environ import *
|
||||
from pyomo.opt import SolverFactory
|
||||
from pyomo.opt import SolverFactory, TerminationCondition, SolverStatus
|
||||
import urllib.parse
|
||||
|
||||
from rota.workers import Worker
|
||||
@@ -211,7 +212,7 @@ class RotaBuilder(object):
|
||||
|
||||
# self.night_blocks = ["weekday", "weekend", "none"]
|
||||
|
||||
self.sites = None
|
||||
self.sites: set = set()
|
||||
|
||||
self.balance_offset_modifier = balance_offset_modifier
|
||||
self.ltft_balance_offset = ltft_balance_offset
|
||||
@@ -246,6 +247,12 @@ class RotaBuilder(object):
|
||||
"avoid_shifts_by_worker_names": [],
|
||||
}
|
||||
|
||||
self.terminate_on_warning = [
|
||||
"Worker/duplicate id",
|
||||
"Worker/duplicate name",
|
||||
"Worker/no valid shifts",
|
||||
]
|
||||
|
||||
self.results = None
|
||||
|
||||
self.warnings: list[tuple[str, str]] = []
|
||||
@@ -294,13 +301,14 @@ class RotaBuilder(object):
|
||||
self.unavailable_to_work_reason = {}
|
||||
self.pref_not_to_work = {}
|
||||
self.pref_not_to_work_reason = {}
|
||||
self.work_requests: set[Tuple[str, WeekInt, DayStr, ShiftName]] = set()
|
||||
self.work_requests: set[
|
||||
Tuple[str | uuid.UUID | int, WeekInt, DayStr, ShiftName]
|
||||
] = set()
|
||||
self.work_requests_map = {}
|
||||
|
||||
def solve_model(
|
||||
self,
|
||||
solver: str = "cbc",
|
||||
use_neos: bool = False,
|
||||
solver: str = "appsi_highs",
|
||||
options: dict = {},
|
||||
debug_if_fail: bool = False,
|
||||
):
|
||||
@@ -314,41 +322,46 @@ class RotaBuilder(object):
|
||||
if "threads" in options:
|
||||
options.pop("threads")
|
||||
self.opt = SolverFactory(solver, executable="scip")
|
||||
elif solver == "appsi_highs":
|
||||
self.opt = SolverFactory(solver)
|
||||
if "seconds" in options:
|
||||
options["time_limit"] = options.pop("seconds")
|
||||
if "ratio" in options:
|
||||
options["mip_rel_gap"] = options.pop("ratio")
|
||||
else:
|
||||
self.opt = SolverFactory(solver)
|
||||
|
||||
try:
|
||||
console.print("Solving")
|
||||
console.print(f"Options: {options}")
|
||||
|
||||
log_file = f"logs/{self.name}_{datetime.datetime.now().strftime('%Y%m%d-%H%M%S')}.log"
|
||||
if use_neos:
|
||||
solver_manager = SolverManagerFactory("neos") # Solve using neos server
|
||||
# results = solver_manager.solve(Rota.model, opt=opt, logfile="test.log")
|
||||
results = solver_manager.solve(
|
||||
self.model,
|
||||
keepfiles=True,
|
||||
tee=True,
|
||||
opt=self.opt,
|
||||
logfile=log_file,
|
||||
)
|
||||
else:
|
||||
results = self.opt.solve(
|
||||
self.model,
|
||||
tee=True,
|
||||
options=options,
|
||||
# options={
|
||||
# "threads": 10,
|
||||
# },
|
||||
logfile=log_file,
|
||||
keepfiles=True,
|
||||
)
|
||||
results = self.opt.solve(
|
||||
self.model,
|
||||
tee=True,
|
||||
options=options,
|
||||
# options={
|
||||
# "threads": 10,
|
||||
# },
|
||||
# logfile=log_file,
|
||||
keepfiles=True,
|
||||
load_solutions=False,
|
||||
)
|
||||
except KeyboardInterrupt:
|
||||
return
|
||||
pass
|
||||
|
||||
if (
|
||||
not results.solver.termination_condition == TerminationCondition.infeasible
|
||||
|
||||
#and not results.solver.status == SolverStatus.aborted
|
||||
):
|
||||
self.model.solutions.load_from(results)
|
||||
|
||||
self.results = results
|
||||
|
||||
console.print(f"Complete - outcome: {results.solver.status}")
|
||||
console.print(f"Termination condition: {results.solver.termination_condition}")
|
||||
|
||||
if results.solver.status != "ok" and debug_if_fail:
|
||||
console.print(f"Attempting each shift individually")
|
||||
@@ -417,7 +430,7 @@ class RotaBuilder(object):
|
||||
solve=True,
|
||||
debug_if_fail: bool = False,
|
||||
export=False,
|
||||
solver="cbc",
|
||||
solver="appsi_highs",
|
||||
):
|
||||
self.run_start_time = datetime.datetime.now()
|
||||
|
||||
@@ -1305,7 +1318,15 @@ class RotaBuilder(object):
|
||||
if self.use_shift_balance_extra:
|
||||
if shift.name in worker.shift_balance_extra:
|
||||
extra = worker.shift_balance_extra[shift.name]
|
||||
target_shifts = target_shifts + extra
|
||||
|
||||
# TODO look at how this affects allocation (how does it affect fte)
|
||||
match extra:
|
||||
case "double":
|
||||
target_shifts = target_shifts * 2
|
||||
case "half":
|
||||
target_shifts = target_shifts / 2
|
||||
case _:
|
||||
target_shifts = target_shifts + extra
|
||||
|
||||
# print(worker.name, shift.name, target_shifts)
|
||||
worker.shift_target_number[shift.name] = target_shifts
|
||||
@@ -2461,6 +2482,15 @@ class RotaBuilder(object):
|
||||
print(f"[bold red]WARNING:[/bold red] {warning_type} - {message}")
|
||||
self.warnings.append((warning_type, message))
|
||||
|
||||
if warning_type in self.terminate_on_warning:
|
||||
raise WarningTermination(warning_type)
|
||||
|
||||
def get_warnings(self, warning_type: None | str = None):
|
||||
if warning_type is None:
|
||||
return self.warnings
|
||||
else:
|
||||
return [warning for warning in self.warnings if warning[0] == warning_type]
|
||||
|
||||
def add_worker(self, worker: Worker) -> None:
|
||||
"""Add a worker to the rota
|
||||
|
||||
@@ -2491,22 +2521,23 @@ class RotaBuilder(object):
|
||||
self.workers_name_map: dict[str, Worker] = {}
|
||||
|
||||
for worker in track(self.workers, description="Building workers"):
|
||||
print(worker.name)
|
||||
worker.load_rota(self)
|
||||
wid = worker.id
|
||||
if wid in self.workers_id_map:
|
||||
raise ValueError(f"Worker with id '{wid}' has been added twice")
|
||||
message = f"Worker with id '{wid}' has been added twice"
|
||||
self.add_warning("Worker/duplicate id", message)
|
||||
|
||||
self.workers_id_map[wid] = worker
|
||||
|
||||
if worker.name in self.workers_name_map:
|
||||
raise ValueError(
|
||||
f"Worker with name '{worker.name}' has been added twice"
|
||||
)
|
||||
message = f"Worker with name '{worker.name}' has been added twice"
|
||||
self.add_warning("Worker/duplicate name", message)
|
||||
self.workers_name_map[worker.name] = worker
|
||||
|
||||
if worker.site not in self.sites:
|
||||
raise ValueError(
|
||||
f"Worker with name '{worker.name}' ({worker.id}) has no valid shifts (site: {worker.site})"
|
||||
)
|
||||
message = f"Worker with name '{worker.name}' ({worker.id}) has no valid shifts (site: {worker.site})"
|
||||
self.add_warning("Worker/no valid shifts", message)
|
||||
|
||||
self.workers = sorted(self.workers)
|
||||
|
||||
@@ -3128,7 +3159,7 @@ class RotaBuilder(object):
|
||||
# Loop through all the days possible shifts and see
|
||||
# if the worker has been assigned
|
||||
for shift in self.get_shift_names_by_week_day(week, day):
|
||||
if model.works[worker.id, week, day, shift].value > 0:
|
||||
if model.works[worker.id, week, day, shift].value > 0.5:
|
||||
shifts.append(shift)
|
||||
a = shift[0]
|
||||
shift_name = shift
|
||||
@@ -3187,6 +3218,7 @@ class RotaBuilder(object):
|
||||
shift_diff_dict[shift.name] = diff
|
||||
|
||||
worker_td = """<td title='Site: {site}' class='worker {site}'
|
||||
data-worker-id='{worker_id}'
|
||||
data-nwds='{nwds}' data-site='{site}' data-worker='{name}'
|
||||
data-fte='{fte}' data-fte_adj='{fte_adj}'
|
||||
data-start_date='{start_date}'
|
||||
@@ -3204,6 +3236,7 @@ class RotaBuilder(object):
|
||||
<span class='name' title='{name}'>{name}</span> ({grade}) [{fte}]</td>""".format(
|
||||
site=worker.site,
|
||||
nwds=nwds,
|
||||
worker_id=worker.id,
|
||||
name=worker.name,
|
||||
fte=worker.fte,
|
||||
fte_adj=worker.fte_adj,
|
||||
@@ -3465,3 +3498,9 @@ class InvalidShift(Exception):
|
||||
"""Raised when there are no active sites"""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
class WarningTermination(Exception):
|
||||
"""Raised when a warning in the termination group is raised"""
|
||||
|
||||
pass
|
||||
|
||||
Reference in New Issue
Block a user