Coverage for backend/django/core/auxiliary/services/parameter_sweep.py: 90%
289 statements
« prev ^ index » next coverage.py v7.10.7, created at 2026-07-22 05:22 +0000
« prev ^ index » next coverage.py v7.10.7, created at 2026-07-22 05:22 +0000
1"""Parameter sweep domain models and row-generation service functions."""
3from __future__ import annotations
5from decimal import Decimal, ROUND_FLOOR
6from enum import StrEnum
7from itertools import product
8import random
9from typing import Any, Mapping
11from django.db.models import Q
12from django.db import transaction
13from pydantic import BaseModel, ConfigDict, Field, RootModel, StrictInt
14from pydantic import ValidationError as PydanticValidationError
15from pydantic import field_validator, model_validator
16from rest_framework.exceptions import ValidationError
18from core.auxiliary.enums.uiEnums import DisplayType
19from core.auxiliary.models.DataColumn import DataColumn
20from core.auxiliary.models.PropertyValue import PropertyValue
21from core.auxiliary.models.Scenario import (
22 ParameterSweepDefinition,
23 ParameterSweepParameter,
24 Scenario,
25 ScenarioInputModeEnum,
26)
27from core.auxiliary.services.mss_input_data import (
28 MssInputCell,
29 MssInputRow,
30 MssInputSource,
31 clear_mss_input_rows,
32 replace_mss_input_rows,
33)
36ROW_WARNING_THRESHOLD = 1000
37ROW_HARD_LIMIT = 50000
38CELL_HARD_LIMIT = 1000000
39DECIMAL_QUANT = Decimal("0.000000000001")
42class ParameterSweepRequestMethodEnum(StrEnum):
43 Grid = "grid"
44 MonteCarlo = "monte_carlo"
45 Hammersley = "hammersley"
46 HaltonZaremba = "halton_zaremba"
49class ParameterSweepParameterRequest(BaseModel):
50 """One free parameter and the numeric interval to explore."""
52 model_config = ConfigDict(extra="forbid")
54 property_value: StrictInt
55 lower_bound: Decimal
56 upper_bound: Decimal
57 step: Decimal | None = None
59 @field_validator("lower_bound", "upper_bound", "step")
60 @classmethod
61 def validate_decimal(cls, value: Decimal | None) -> Decimal | None:
62 if value is not None and not value.is_finite(): 62 ↛ 63line 62 didn't jump to line 63 because the condition on line 62 was never true
63 raise ValueError("Numeric values must be finite.")
64 return value
67class ParameterSweepRequest(BaseModel):
68 """Validated parameter sweep request shared by preview and generation."""
70 model_config = ConfigDict(extra="forbid")
72 method: ParameterSweepRequestMethodEnum = ParameterSweepRequestMethodEnum.Grid
73 parameters: list[ParameterSweepParameterRequest] = Field(min_length=1)
74 sample_count: StrictInt | None = None
75 monte_carlo_seed: StrictInt | None = None
76 confirm_replace: bool = False
77 confirm_large_generation: bool = False
79 @model_validator(mode="after")
80 def validate_method_requirements(self) -> "ParameterSweepRequest":
81 if self.method == ParameterSweepRequestMethodEnum.Grid:
82 self.sample_count = None
83 self.monte_carlo_seed = None
84 for parameter in self.parameters:
85 _validate_grid_parameter(parameter)
86 return self
88 if self.sample_count is None: 88 ↛ 89line 88 didn't jump to line 89 because the condition on line 88 was never true
89 raise ValueError("Sample count is required for sampling methods.")
90 if self.sample_count <= 0: 90 ↛ 91line 90 didn't jump to line 91 because the condition on line 90 was never true
91 raise ValueError("Sample count must be greater than zero.")
93 if self.monte_carlo_seed is not None and self.monte_carlo_seed < 0: 93 ↛ 94line 93 didn't jump to line 94 because the condition on line 93 was never true
94 raise ValueError("Seed must be a non-negative integer.")
95 if self.method == ParameterSweepRequestMethodEnum.MonteCarlo:
96 if self.monte_carlo_seed is None:
97 self.monte_carlo_seed = random.SystemRandom().randrange(1, 2**31)
98 else:
99 self.monte_carlo_seed = None
100 for parameter in self.parameters:
101 parameter.step = None
102 return self
105class EligibleParameterSweepTarget(BaseModel):
106 property_value: int
107 simulation_object: int
108 simulation_object_name: str
109 property_name: str
110 indexed_set_names: list[str]
111 unit: str
112 label: str
115class EligibleParameterSweepTargetsResponse(RootModel[list[EligibleParameterSweepTarget]]):
116 """List response wrapper so OpenAPI keeps the target endpoint typed as an array."""
119class ParameterSweepPreviewParameterResponse(BaseModel):
120 property_value: int
121 label: str
122 unit: str
123 value_count: int
126class ParameterSweepPreviewResponse(BaseModel):
127 method: ParameterSweepRequestMethodEnum
128 row_count: int
129 warns_above_threshold: bool
130 hard_limit: int
131 warning_threshold: int
132 monte_carlo_seed: int | None = None
133 parameters: list[ParameterSweepPreviewParameterResponse]
136class ParameterSweepGenerateResponse(BaseModel):
137 method: ParameterSweepRequestMethodEnum
138 row_count: int
139 definition: int
140 monte_carlo_seed: int | None = None
143def eligible_parameter_sweep_targets(flowsheet_id: int) -> list[EligibleParameterSweepTarget]:
144 """Return directly settable numeric property values for a flowsheet sweep."""
146 values = (
147 PropertyValue.objects.filter(
148 flowsheet_state__flowsheet_id=flowsheet_id,
149 property__type=DisplayType.numeric,
150 property__set__simulationObject__is_deleted=False,
151 )
152 .filter(Q(formula__isnull=True) | Q(formula=""))
153 .filter(Q(enabled=True) | Q(controlSetPoint__isnull=False))
154 .filter(controlManipulated__isnull=True)
155 .select_related("property", "property__set", "property__set__simulationObject")
156 .prefetch_related("indexedItems")
157 .order_by(
158 "property__set__simulationObject__componentName",
159 "property__displayName",
160 "id",
161 )
162 )
164 return [_serialize_target(value) for value in values]
167def preview_parameter_sweep(
168 scenario: Scenario,
169 payload: ParameterSweepRequest | Mapping[str, Any],
170) -> ParameterSweepPreviewResponse:
171 """Validate a sweep request and report its size without persisting rows."""
173 spec = _coerce_sweep_request(payload)
174 targets = _validate_targets(scenario, spec)
175 row_count = _calculate_row_count(spec)
176 _validate_total_cell_count(row_count, len(targets))
177 return ParameterSweepPreviewResponse(
178 method=spec.method,
179 row_count=row_count,
180 warns_above_threshold=row_count > ROW_WARNING_THRESHOLD,
181 hard_limit=ROW_HARD_LIMIT,
182 warning_threshold=ROW_WARNING_THRESHOLD,
183 monte_carlo_seed=spec.monte_carlo_seed,
184 parameters=[
185 ParameterSweepPreviewParameterResponse(
186 property_value=target.id,
187 label=_target_label(target),
188 unit=target.property.unit,
189 value_count=_parameter_value_count(spec, index),
190 )
191 for index, target in enumerate(targets)
192 ],
193 )
196@transaction.atomic
197def generate_parameter_sweep(
198 scenario: Scenario,
199 payload: ParameterSweepRequest | Mapping[str, Any],
200) -> ParameterSweepGenerateResponse:
201 """Replace a scenario's MSS input with generated sweep rows.
203 The operation is deliberately all-or-nothing because the scenario mode,
204 saved sweep definition, data columns, rows, and cells must remain in sync.
205 """
207 spec = _coerce_sweep_request(payload)
208 targets = _validate_targets(scenario, spec)
209 row_count = _calculate_row_count(spec)
210 _validate_total_cell_count(row_count, len(targets))
211 existing_rows = scenario.dataRows.exists()
212 existing_columns = scenario.dataColumns.exists()
213 requires_replace = existing_rows or existing_columns
215 if requires_replace and not spec.confirm_replace:
216 raise ValidationError(
217 {
218 "confirm_replace": (
219 "Existing scenario input data will be replaced. "
220 "Set confirm_replace to true to continue."
221 )
222 }
223 )
224 if row_count > ROW_WARNING_THRESHOLD and not spec.confirm_large_generation:
225 raise ValidationError(
226 {
227 "confirm_large_generation": (
228 f"This sweep will generate {row_count} rows. "
229 "Set confirm_large_generation to true to continue."
230 )
231 }
232 )
234 rows = _generate_rows(spec)
236 DataColumn.objects.filter(scenario=scenario).delete()
237 ParameterSweepDefinition.objects.filter(scenario=scenario).delete()
239 definition = ParameterSweepDefinition.objects.create(
240 flowsheet_state=scenario.flowsheet_state,
241 scenario=scenario,
242 method=spec.method,
243 sample_count=spec.sample_count,
244 monte_carlo_seed=spec.monte_carlo_seed,
245 )
246 ParameterSweepParameter.objects.bulk_create(
247 [
248 ParameterSweepParameter(
249 flowsheet_state=scenario.flowsheet_state,
250 definition=definition,
251 property_value=target,
252 order=index,
253 lower_bound=param.lower_bound,
254 upper_bound=param.upper_bound,
255 step=param.step,
256 unit=target.property.unit,
257 target_label=_target_label(target),
258 )
259 for index, (param, target) in enumerate(zip(spec.parameters, targets))
260 ]
261 )
263 columns = [
264 DataColumn(
265 flowsheet_state=scenario.flowsheet_state,
266 scenario=scenario,
267 name=_unique_column_name(target, index, targets),
268 property_value=target,
269 )
270 for index, target in enumerate(targets)
271 ]
272 DataColumn.objects.bulk_create(columns)
273 columns = list(DataColumn.objects.filter(scenario=scenario).order_by("id"))
275 replace_mss_input_rows(
276 scenario=scenario,
277 rows=[
278 MssInputRow(
279 index=index,
280 cells=[
281 MssInputCell(
282 data_column_id=column.id,
283 value=float(value),
284 )
285 for column, value in zip(columns, values)
286 ],
287 )
288 for index, values in enumerate(rows)
289 ],
290 source=MssInputSource.parameter_sweep,
291 )
293 scenario.mss_input_mode = ScenarioInputModeEnum.ParameterSweep
294 scenario.Uploaded_fileName = ""
295 scenario.save(update_fields=["mss_input_mode", "Uploaded_fileName"])
297 return ParameterSweepGenerateResponse(
298 method=spec.method,
299 row_count=row_count,
300 definition=definition.id,
301 monte_carlo_seed=spec.monte_carlo_seed,
302 )
305def clear_parameter_sweep_definition(scenario: Scenario) -> None:
306 ParameterSweepDefinition.objects.filter(scenario=scenario).delete()
309def clear_mss_input_data(
310 scenario: Scenario,
311 *,
312 clear_uploaded_filename: bool = True,
313) -> None:
314 """Remove all generated/uploaded MSS table data for a scenario."""
316 clear_mss_input_rows(scenario=scenario, source=MssInputSource.mode_switch)
317 DataColumn.objects.filter(scenario=scenario).delete()
318 clear_parameter_sweep_definition(scenario)
319 if clear_uploaded_filename: 319 ↛ exitline 319 didn't return from function 'clear_mss_input_data' because the condition on line 319 was always true
320 scenario.Uploaded_fileName = ""
321 scenario.save(update_fields=["Uploaded_fileName"])
324def clear_mss_input_data_after_mode_switch(
325 scenario: Scenario,
326 *,
327 previous_mode: str,
328 requested_mode: str,
329 had_parameter_sweep_definition: bool,
330) -> None:
331 """Clear stale MSS data after the scenario's persisted input mode changes."""
333 if requested_mode not in { 333 ↛ 337line 333 didn't jump to line 337 because the condition on line 333 was never true
334 ScenarioInputModeEnum.Csv,
335 ScenarioInputModeEnum.ParameterSweep,
336 }:
337 return
339 if previous_mode == requested_mode and not ( 339 ↛ 342line 339 didn't jump to line 342 because the condition on line 339 was never true
340 requested_mode == ScenarioInputModeEnum.Csv and had_parameter_sweep_definition
341 ):
342 return
344 clear_mss_input_data(scenario)
347def validate_parameter_sweep_solve_ready(scenario: Scenario) -> None:
348 if scenario.mss_input_mode != ScenarioInputModeEnum.ParameterSweep:
349 return
351 try:
352 definition = scenario.parameterSweepDefinition
353 except ParameterSweepDefinition.DoesNotExist as exc:
354 raise ValidationError("Parameter sweep scenarios require a saved sweep definition.") from exc
356 params = list(definition.parameters.select_related("property_value"))
357 if not params: 357 ↛ 358line 357 didn't jump to line 358 because the condition on line 357 was never true
358 raise ValidationError("Parameter sweep scenarios require at least one parameter.")
359 if not scenario.dataRows.exists(): 359 ↛ 360line 359 didn't jump to line 360 because the condition on line 359 was never true
360 raise ValidationError("Parameter sweep scenarios require generated rows.")
362 spec = ParameterSweepRequest(
363 method=definition.method,
364 sample_count=definition.sample_count,
365 monte_carlo_seed=definition.monte_carlo_seed,
366 parameters=[
367 ParameterSweepParameterRequest(
368 property_value=param.property_value_id,
369 lower_bound=param.lower_bound,
370 upper_bound=param.upper_bound,
371 step=param.step,
372 )
373 for param in params
374 ],
375 )
376 _validate_targets(scenario, spec)
379def _generate_rows(spec: ParameterSweepRequest) -> list[list[Decimal]]:
380 """Materialize the bounded samples in the column order chosen by the user."""
382 if spec.method == ParameterSweepRequestMethodEnum.Grid:
383 value_lists = [_grid_values(parameter) for parameter in spec.parameters]
384 rows = [list(values) for values in product(*value_lists)]
385 elif spec.method == ParameterSweepRequestMethodEnum.MonteCarlo:
386 rows = _monte_carlo_rows(spec)
387 elif spec.method == ParameterSweepRequestMethodEnum.Hammersley:
388 rows = _hammersley_rows(spec)
389 elif spec.method == ParameterSweepRequestMethodEnum.HaltonZaremba: 389 ↛ 392line 389 didn't jump to line 392 because the condition on line 389 was always true
390 rows = _halton_rows(spec)
391 else:
392 raise ValidationError({"method": f"Unsupported parameter sweep method: {spec.method}"})
394 return rows
397def _calculate_row_count(spec: ParameterSweepRequest) -> int:
398 """Calculate sweep size before materializing rows so hard limits stay cheap."""
400 if spec.method != ParameterSweepRequestMethodEnum.Grid:
401 sample_count = spec.sample_count or 0
402 if sample_count > ROW_HARD_LIMIT:
403 raise ValidationError(
404 {
405 "row_count": (
406 f"Parameter sweep cannot generate more than {ROW_HARD_LIMIT} rows."
407 )
408 }
409 )
410 return sample_count
412 total = 1
413 for parameter in spec.parameters:
414 total *= _grid_value_count(parameter)
415 if total > ROW_HARD_LIMIT:
416 raise ValidationError(
417 {
418 "row_count": (
419 f"Parameter sweep cannot generate more than {ROW_HARD_LIMIT} rows."
420 )
421 }
422 )
423 return total
426def _validate_total_cell_count(row_count: int, parameter_count: int) -> None:
427 total_cells = row_count * parameter_count
428 if total_cells > CELL_HARD_LIMIT:
429 raise ValidationError(
430 {
431 "cell_count": (
432 f"Parameter sweep cannot generate more than {CELL_HARD_LIMIT} cells."
433 )
434 }
435 )
438def _grid_values(parameter: ParameterSweepParameterRequest) -> list[Decimal]:
439 count = _grid_value_count(parameter)
440 current = parameter.lower_bound
441 values: list[Decimal] = []
442 for _ in range(count):
443 values.append(current.quantize(DECIMAL_QUANT))
444 current += parameter.step or 0
445 return values
448def _grid_value_count(parameter: ParameterSweepParameterRequest) -> int:
449 step = parameter.step
450 if step is None: 450 ↛ 451line 450 didn't jump to line 451 because the condition on line 450 was never true
451 raise ValidationError({"step": "Grid sweeps require a step for every parameter."})
452 if step == 0: 452 ↛ 453line 452 didn't jump to line 453 because the condition on line 452 was never true
453 raise ValidationError({"step": "Step cannot be zero."})
455 start = parameter.lower_bound
456 end = parameter.upper_bound
457 if start == end:
458 return 1
460 if (end > start and step < 0) or (end < start and step > 0): 460 ↛ 461line 460 didn't jump to line 461 because the condition on line 460 was never true
461 raise ValidationError({"step": "Step must move from start toward end."})
463 span = abs(end - start)
464 step_size = abs(step)
465 return int((span / step_size).to_integral_value(rounding=ROUND_FLOOR)) + 1
468def _monte_carlo_rows(spec: ParameterSweepRequest) -> list[list[Decimal]]:
469 rng = random.Random(spec.monte_carlo_seed)
470 return [
471 [
472 _scale_unit_interval(Decimal(str(rng.random())), parameter)
473 for parameter in spec.parameters
474 ]
475 for _ in range(spec.sample_count or 0)
476 ]
479def _hammersley_rows(spec: ParameterSweepRequest) -> list[list[Decimal]]:
480 from skopt.sampler import Hammersly
481 from skopt.space import Real
483 dimensions = [Real(0.0, 1.0) for _ in spec.parameters]
484 points = Hammersly().generate(dimensions, spec.sample_count or 0)
485 return _scale_points(points, spec.parameters)
488def _halton_rows(spec: ParameterSweepRequest) -> list[list[Decimal]]:
489 from scipy.stats import qmc
491 sampler = qmc.Halton(d=len(spec.parameters), scramble=False)
492 points = sampler.random(spec.sample_count or 0)
493 return _scale_points(points, spec.parameters)
496def _scale_points(
497 points,
498 parameters: list[ParameterSweepParameterRequest],
499) -> list[list[Decimal]]:
500 return [
501 [
502 _scale_unit_interval(Decimal(str(point[index])), parameter)
503 for index, parameter in enumerate(parameters)
504 ]
505 for point in points
506 ]
509def _scale_unit_interval(
510 value: Decimal,
511 parameter: ParameterSweepParameterRequest,
512) -> Decimal:
513 lower = parameter.lower_bound
514 upper = parameter.upper_bound
515 return (lower + (upper - lower) * value).quantize(DECIMAL_QUANT)
518def _validate_targets(
519 scenario: Scenario,
520 spec: ParameterSweepRequest,
521) -> list[PropertyValue]:
522 """Ensure every requested target is still a settable numeric variable."""
524 target_ids = [param.property_value for param in spec.parameters]
525 if len(set(target_ids)) != len(target_ids): 525 ↛ 526line 525 didn't jump to line 526 because the condition on line 525 was never true
526 raise ValidationError({"parameters": "Each parameter can only be selected once."})
528 targets = list(
529 PropertyValue.objects.filter(id__in=target_ids)
530 .select_related("property", "property__set", "property__set__simulationObject")
531 .prefetch_related("indexedItems")
532 )
533 targets_by_id = {target.id: target for target in targets}
534 ordered_targets: list[PropertyValue] = []
536 for target_id in target_ids:
537 target = targets_by_id.get(target_id)
538 if target is None: 538 ↛ 539line 538 didn't jump to line 539 because the condition on line 538 was never true
539 raise ValidationError({"parameters": f"Property value {target_id} was not found."})
540 if target.flowsheet_state_id != scenario.flowsheet_state_id: 540 ↛ 541line 540 didn't jump to line 541 because the condition on line 540 was never true
541 raise ValidationError({"parameters": "Sweep targets must belong to the scenario flowsheet."})
542 if not _is_eligible_target(target): 542 ↛ 543line 542 didn't jump to line 543 because the condition on line 542 was never true
543 raise ValidationError(
544 {
545 "parameters": (
546 f"{_target_label(target)} is no longer eligible for parameter sweep."
547 )
548 }
549 )
550 ordered_targets.append(target)
552 return ordered_targets
555def _is_eligible_target(value: PropertyValue) -> bool:
556 return (
557 value.is_enabled()
558 and not value.formula
559 and value.property is not None
560 and value.property.type == DisplayType.numeric
561 and not hasattr(value, "controlManipulated")
562 and value.property.set is not None
563 and value.property.set.simulationObject is not None
564 and not value.property.set.simulationObject.is_deleted
565 )
568def _parameter_value_count(spec: ParameterSweepRequest, index: int) -> int:
569 if spec.method == ParameterSweepRequestMethodEnum.Grid:
570 return _grid_value_count(spec.parameters[index])
571 return spec.sample_count or 0
574def _serialize_target(value: PropertyValue) -> EligibleParameterSweepTarget:
575 return EligibleParameterSweepTarget(
576 property_value=value.id,
577 simulation_object=value.property.set.simulationObject.id,
578 simulation_object_name=value.property.set.simulationObject.componentName,
579 property_name=value.property.displayName,
580 indexed_set_names=value.get_index_names(),
581 unit=value.property.unit,
582 label=_target_label(value),
583 )
586def _target_label(value: PropertyValue) -> str:
587 object_name = value.property.set.simulationObject.componentName or "Object"
588 index_names = value.get_index_names()
589 suffix = f" {' '.join(index_names)}" if index_names else ""
590 return f"{object_name} / {value.property.displayName}{suffix}"
593def _unique_column_name(target: PropertyValue, index: int, targets: list[PropertyValue]) -> str:
594 label = _target_label(target)
595 if sum(1 for candidate in targets if _target_label(candidate) == label) == 1: 595 ↛ 597line 595 didn't jump to line 597 because the condition on line 595 was always true
596 return label
597 return f"{label} ({target.id})"
600def _coerce_sweep_request(
601 payload: ParameterSweepRequest | Mapping[str, Any],
602) -> ParameterSweepRequest:
603 if isinstance(payload, ParameterSweepRequest):
604 return payload
606 # The frontend base query injects the active flowsheet id into mutation
607 # bodies as transport context. It is not part of the sweep contract, so
608 # remove it before applying the strict Pydantic request model.
609 payload_without_transport_context = dict(payload)
610 payload_without_transport_context.pop("flowsheet", None)
612 try:
613 return ParameterSweepRequest.model_validate(payload_without_transport_context)
614 except PydanticValidationError as exc:
615 raise _drf_validation_error_from_pydantic(exc) from exc
618def _validate_grid_parameter(parameter: ParameterSweepParameterRequest) -> None:
619 if parameter.step is None: 619 ↛ 620line 619 didn't jump to line 620 because the condition on line 619 was never true
620 raise ValueError("Grid sweeps require a step for every parameter.")
621 if parameter.step == 0: 621 ↛ 622line 621 didn't jump to line 622 because the condition on line 621 was never true
622 raise ValueError("Step cannot be zero.")
624 if parameter.lower_bound == parameter.upper_bound:
625 return
626 if (
627 parameter.upper_bound > parameter.lower_bound
628 and parameter.step < 0
629 ) or (
630 parameter.upper_bound < parameter.lower_bound
631 and parameter.step > 0
632 ):
633 raise ValueError("Step must move from start toward end.")
636def _drf_validation_error_from_pydantic(exc: PydanticValidationError) -> ValidationError:
637 details: dict[str, list[str]] = {}
638 for error in exc.errors():
639 location = error.get("loc") or ("non_field_errors",)
640 field = str(location[0])
641 details.setdefault(field, []).append(str(error.get("msg", "Invalid value.")))
642 return ValidationError(details)