update from last round

This commit is contained in:
Ross
2024-12-23 16:56:53 +00:00
parent 51b4e89601
commit 70dc279a02
13 changed files with 161 additions and 100 deletions
+76 -37
View File
@@ -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