phased-array-systems 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- phased_array_systems/__about__.py +4 -0
- phased_array_systems/__init__.py +10 -0
- phased_array_systems/architecture/__init__.py +15 -0
- phased_array_systems/architecture/config.py +152 -0
- phased_array_systems/cli.py +25 -0
- phased_array_systems/constants.py +55 -0
- phased_array_systems/evaluate.py +136 -0
- phased_array_systems/io/__init__.py +13 -0
- phased_array_systems/io/config_loader.py +86 -0
- phased_array_systems/io/exporters.py +171 -0
- phased_array_systems/io/schema.py +145 -0
- phased_array_systems/models/__init__.py +5 -0
- phased_array_systems/models/antenna/__init__.py +15 -0
- phased_array_systems/models/antenna/adapter.py +190 -0
- phased_array_systems/models/antenna/metrics.py +166 -0
- phased_array_systems/models/base.py +30 -0
- phased_array_systems/models/comms/__init__.py +9 -0
- phased_array_systems/models/comms/link_budget.py +171 -0
- phased_array_systems/models/comms/propagation.py +84 -0
- phased_array_systems/models/swapc/__init__.py +9 -0
- phased_array_systems/models/swapc/cost.py +98 -0
- phased_array_systems/models/swapc/power.py +102 -0
- phased_array_systems/requirements/__init__.py +15 -0
- phased_array_systems/requirements/core.py +244 -0
- phased_array_systems/scenarios/__init__.py +11 -0
- phased_array_systems/scenarios/base.py +30 -0
- phased_array_systems/scenarios/comms.py +56 -0
- phased_array_systems/scenarios/radar.py +42 -0
- phased_array_systems/trades/__init__.py +16 -0
- phased_array_systems/trades/design_space.py +241 -0
- phased_array_systems/trades/doe.py +146 -0
- phased_array_systems/trades/pareto.py +266 -0
- phased_array_systems/trades/runner.py +245 -0
- phased_array_systems/types.py +54 -0
- phased_array_systems/utils/__init__.py +8 -0
- phased_array_systems/utils/hashing.py +70 -0
- phased_array_systems/viz/__init__.py +9 -0
- phased_array_systems/viz/plots.py +324 -0
- phased_array_systems-0.1.0.dist-info/METADATA +174 -0
- phased_array_systems-0.1.0.dist-info/RECORD +43 -0
- phased_array_systems-0.1.0.dist-info/WHEEL +4 -0
- phased_array_systems-0.1.0.dist-info/entry_points.txt +2 -0
- phased_array_systems-0.1.0.dist-info/licenses/LICENSE +21 -0
|
@@ -0,0 +1,244 @@
|
|
|
1
|
+
"""Core requirement and verification classes."""
|
|
2
|
+
|
|
3
|
+
from dataclasses import dataclass, field
|
|
4
|
+
|
|
5
|
+
from phased_array_systems.types import ComparisonOp, MetricsDict, Severity
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
@dataclass(frozen=True)
|
|
9
|
+
class Requirement:
|
|
10
|
+
"""A single requirement specification.
|
|
11
|
+
|
|
12
|
+
Attributes:
|
|
13
|
+
id: Unique identifier for the requirement (e.g., "REQ-001")
|
|
14
|
+
name: Human-readable name
|
|
15
|
+
metric_key: The key in the metrics dictionary to check
|
|
16
|
+
op: Comparison operator
|
|
17
|
+
value: Threshold value to compare against
|
|
18
|
+
units: Optional units string for documentation
|
|
19
|
+
severity: Importance level ("must", "should", "nice")
|
|
20
|
+
"""
|
|
21
|
+
|
|
22
|
+
id: str
|
|
23
|
+
name: str
|
|
24
|
+
metric_key: str
|
|
25
|
+
op: ComparisonOp
|
|
26
|
+
value: float
|
|
27
|
+
units: str | None = None
|
|
28
|
+
severity: Severity = "must"
|
|
29
|
+
|
|
30
|
+
def check(self, actual_value: float) -> bool:
|
|
31
|
+
"""Check if the actual value satisfies this requirement.
|
|
32
|
+
|
|
33
|
+
Args:
|
|
34
|
+
actual_value: The measured/computed value to check
|
|
35
|
+
|
|
36
|
+
Returns:
|
|
37
|
+
True if requirement is satisfied, False otherwise
|
|
38
|
+
"""
|
|
39
|
+
if self.op == ">=":
|
|
40
|
+
return actual_value >= self.value
|
|
41
|
+
elif self.op == "<=":
|
|
42
|
+
return actual_value <= self.value
|
|
43
|
+
elif self.op == "==":
|
|
44
|
+
return actual_value == self.value
|
|
45
|
+
elif self.op == ">":
|
|
46
|
+
return actual_value > self.value
|
|
47
|
+
elif self.op == "<":
|
|
48
|
+
return actual_value < self.value
|
|
49
|
+
else:
|
|
50
|
+
raise ValueError(f"Unknown operator: {self.op}")
|
|
51
|
+
|
|
52
|
+
def compute_margin(self, actual_value: float) -> float:
|
|
53
|
+
"""Compute the margin to the requirement threshold.
|
|
54
|
+
|
|
55
|
+
Positive margin means the requirement is satisfied with room to spare.
|
|
56
|
+
Negative margin means the requirement is not met.
|
|
57
|
+
|
|
58
|
+
Args:
|
|
59
|
+
actual_value: The measured/computed value
|
|
60
|
+
|
|
61
|
+
Returns:
|
|
62
|
+
Margin value (interpretation depends on operator)
|
|
63
|
+
"""
|
|
64
|
+
if self.op in (">=", ">"):
|
|
65
|
+
return actual_value - self.value
|
|
66
|
+
elif self.op in ("<=", "<"):
|
|
67
|
+
return self.value - actual_value
|
|
68
|
+
else: # ==
|
|
69
|
+
return -abs(actual_value - self.value)
|
|
70
|
+
|
|
71
|
+
|
|
72
|
+
@dataclass
|
|
73
|
+
class RequirementResult:
|
|
74
|
+
"""Result of checking a single requirement.
|
|
75
|
+
|
|
76
|
+
Attributes:
|
|
77
|
+
requirement: The requirement that was checked
|
|
78
|
+
actual_value: The actual value from metrics (None if metric missing)
|
|
79
|
+
passes: Whether the requirement passed
|
|
80
|
+
margin: Margin to the threshold
|
|
81
|
+
error: Error message if metric was missing or check failed
|
|
82
|
+
"""
|
|
83
|
+
|
|
84
|
+
requirement: Requirement
|
|
85
|
+
actual_value: float | None
|
|
86
|
+
passes: bool
|
|
87
|
+
margin: float | None
|
|
88
|
+
error: str | None = None
|
|
89
|
+
|
|
90
|
+
|
|
91
|
+
@dataclass
|
|
92
|
+
class VerificationReport:
|
|
93
|
+
"""Complete verification report for a set of requirements.
|
|
94
|
+
|
|
95
|
+
Attributes:
|
|
96
|
+
passes: True if ALL 'must' requirements pass
|
|
97
|
+
results: List of individual requirement results
|
|
98
|
+
failed_ids: List of requirement IDs that failed
|
|
99
|
+
must_pass_count: Number of must requirements that passed
|
|
100
|
+
must_total_count: Total number of must requirements
|
|
101
|
+
should_pass_count: Number of should requirements that passed
|
|
102
|
+
should_total_count: Total number of should requirements
|
|
103
|
+
"""
|
|
104
|
+
|
|
105
|
+
passes: bool
|
|
106
|
+
results: list[RequirementResult]
|
|
107
|
+
failed_ids: list[str]
|
|
108
|
+
must_pass_count: int = 0
|
|
109
|
+
must_total_count: int = 0
|
|
110
|
+
should_pass_count: int = 0
|
|
111
|
+
should_total_count: int = 0
|
|
112
|
+
|
|
113
|
+
def to_dict(self) -> dict:
|
|
114
|
+
"""Convert report to dictionary for serialization."""
|
|
115
|
+
return {
|
|
116
|
+
"passes": self.passes,
|
|
117
|
+
"failed_ids": self.failed_ids,
|
|
118
|
+
"must_pass_count": self.must_pass_count,
|
|
119
|
+
"must_total_count": self.must_total_count,
|
|
120
|
+
"should_pass_count": self.should_pass_count,
|
|
121
|
+
"should_total_count": self.should_total_count,
|
|
122
|
+
"results": [
|
|
123
|
+
{
|
|
124
|
+
"id": r.requirement.id,
|
|
125
|
+
"name": r.requirement.name,
|
|
126
|
+
"metric_key": r.requirement.metric_key,
|
|
127
|
+
"threshold": r.requirement.value,
|
|
128
|
+
"operator": r.requirement.op,
|
|
129
|
+
"actual_value": r.actual_value,
|
|
130
|
+
"passes": r.passes,
|
|
131
|
+
"margin": r.margin,
|
|
132
|
+
"severity": r.requirement.severity,
|
|
133
|
+
"error": r.error,
|
|
134
|
+
}
|
|
135
|
+
for r in self.results
|
|
136
|
+
],
|
|
137
|
+
}
|
|
138
|
+
|
|
139
|
+
|
|
140
|
+
@dataclass
|
|
141
|
+
class RequirementSet:
|
|
142
|
+
"""A collection of requirements with verification capabilities.
|
|
143
|
+
|
|
144
|
+
Attributes:
|
|
145
|
+
requirements: List of requirements
|
|
146
|
+
name: Optional name for the requirement set
|
|
147
|
+
"""
|
|
148
|
+
|
|
149
|
+
requirements: list[Requirement] = field(default_factory=list)
|
|
150
|
+
name: str | None = None
|
|
151
|
+
|
|
152
|
+
def add(self, requirement: Requirement) -> None:
|
|
153
|
+
"""Add a requirement to the set."""
|
|
154
|
+
self.requirements.append(requirement)
|
|
155
|
+
|
|
156
|
+
def verify(self, metrics: MetricsDict) -> VerificationReport:
|
|
157
|
+
"""Verify all requirements against provided metrics.
|
|
158
|
+
|
|
159
|
+
Args:
|
|
160
|
+
metrics: Dictionary of metric_name -> value
|
|
161
|
+
|
|
162
|
+
Returns:
|
|
163
|
+
VerificationReport with pass/fail status and margins
|
|
164
|
+
"""
|
|
165
|
+
results: list[RequirementResult] = []
|
|
166
|
+
failed_ids: list[str] = []
|
|
167
|
+
must_pass = 0
|
|
168
|
+
must_total = 0
|
|
169
|
+
should_pass = 0
|
|
170
|
+
should_total = 0
|
|
171
|
+
|
|
172
|
+
for req in self.requirements:
|
|
173
|
+
# Track totals by severity
|
|
174
|
+
if req.severity == "must":
|
|
175
|
+
must_total += 1
|
|
176
|
+
elif req.severity == "should":
|
|
177
|
+
should_total += 1
|
|
178
|
+
|
|
179
|
+
# Check if metric exists
|
|
180
|
+
if req.metric_key not in metrics:
|
|
181
|
+
result = RequirementResult(
|
|
182
|
+
requirement=req,
|
|
183
|
+
actual_value=None,
|
|
184
|
+
passes=False,
|
|
185
|
+
margin=None,
|
|
186
|
+
error=f"Metric '{req.metric_key}' not found in results",
|
|
187
|
+
)
|
|
188
|
+
failed_ids.append(req.id)
|
|
189
|
+
else:
|
|
190
|
+
actual = metrics[req.metric_key]
|
|
191
|
+
if actual is None or not isinstance(actual, (int, float)):
|
|
192
|
+
result = RequirementResult(
|
|
193
|
+
requirement=req,
|
|
194
|
+
actual_value=None,
|
|
195
|
+
passes=False,
|
|
196
|
+
margin=None,
|
|
197
|
+
error=f"Metric '{req.metric_key}' has invalid value: {actual}",
|
|
198
|
+
)
|
|
199
|
+
failed_ids.append(req.id)
|
|
200
|
+
else:
|
|
201
|
+
actual_float = float(actual)
|
|
202
|
+
passes = req.check(actual_float)
|
|
203
|
+
margin = req.compute_margin(actual_float)
|
|
204
|
+
result = RequirementResult(
|
|
205
|
+
requirement=req,
|
|
206
|
+
actual_value=actual_float,
|
|
207
|
+
passes=passes,
|
|
208
|
+
margin=margin,
|
|
209
|
+
)
|
|
210
|
+
if not passes:
|
|
211
|
+
failed_ids.append(req.id)
|
|
212
|
+
else:
|
|
213
|
+
if req.severity == "must":
|
|
214
|
+
must_pass += 1
|
|
215
|
+
elif req.severity == "should":
|
|
216
|
+
should_pass += 1
|
|
217
|
+
|
|
218
|
+
results.append(result)
|
|
219
|
+
|
|
220
|
+
# Overall pass requires all "must" requirements to pass
|
|
221
|
+
all_must_pass = must_pass == must_total
|
|
222
|
+
|
|
223
|
+
return VerificationReport(
|
|
224
|
+
passes=all_must_pass,
|
|
225
|
+
results=results,
|
|
226
|
+
failed_ids=failed_ids,
|
|
227
|
+
must_pass_count=must_pass,
|
|
228
|
+
must_total_count=must_total,
|
|
229
|
+
should_pass_count=should_pass,
|
|
230
|
+
should_total_count=should_total,
|
|
231
|
+
)
|
|
232
|
+
|
|
233
|
+
def get_by_id(self, req_id: str) -> Requirement | None:
|
|
234
|
+
"""Get a requirement by its ID."""
|
|
235
|
+
for req in self.requirements:
|
|
236
|
+
if req.id == req_id:
|
|
237
|
+
return req
|
|
238
|
+
return None
|
|
239
|
+
|
|
240
|
+
def __len__(self) -> int:
|
|
241
|
+
return len(self.requirements)
|
|
242
|
+
|
|
243
|
+
def __iter__(self):
|
|
244
|
+
return iter(self.requirements)
|
|
@@ -0,0 +1,11 @@
|
|
|
1
|
+
"""Scenario definitions for communications and radar applications."""
|
|
2
|
+
|
|
3
|
+
from phased_array_systems.scenarios.base import ScenarioBase
|
|
4
|
+
from phased_array_systems.scenarios.comms import CommsLinkScenario
|
|
5
|
+
from phased_array_systems.scenarios.radar import RadarDetectionScenario
|
|
6
|
+
|
|
7
|
+
__all__ = [
|
|
8
|
+
"ScenarioBase",
|
|
9
|
+
"CommsLinkScenario",
|
|
10
|
+
"RadarDetectionScenario",
|
|
11
|
+
]
|
|
@@ -0,0 +1,30 @@
|
|
|
1
|
+
"""Base scenario class and utilities."""
|
|
2
|
+
|
|
3
|
+
from pydantic import BaseModel, Field
|
|
4
|
+
|
|
5
|
+
from phased_array_systems.constants import C
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
class ScenarioBase(BaseModel):
|
|
9
|
+
"""Base class for all scenario types.
|
|
10
|
+
|
|
11
|
+
All scenarios must have a frequency, which is used to compute
|
|
12
|
+
wavelength and other frequency-dependent parameters.
|
|
13
|
+
|
|
14
|
+
Attributes:
|
|
15
|
+
freq_hz: Operating frequency in Hz
|
|
16
|
+
name: Optional name for the scenario
|
|
17
|
+
"""
|
|
18
|
+
|
|
19
|
+
freq_hz: float = Field(gt=0, description="Operating frequency (Hz)")
|
|
20
|
+
name: str | None = Field(default=None, description="Scenario name")
|
|
21
|
+
|
|
22
|
+
@property
|
|
23
|
+
def wavelength_m(self) -> float:
|
|
24
|
+
"""Wavelength in meters."""
|
|
25
|
+
return C / self.freq_hz
|
|
26
|
+
|
|
27
|
+
@property
|
|
28
|
+
def freq_ghz(self) -> float:
|
|
29
|
+
"""Frequency in GHz."""
|
|
30
|
+
return self.freq_hz / 1e9
|
|
@@ -0,0 +1,56 @@
|
|
|
1
|
+
"""Communications link scenario definition."""
|
|
2
|
+
|
|
3
|
+
from typing import Literal
|
|
4
|
+
|
|
5
|
+
from pydantic import Field
|
|
6
|
+
|
|
7
|
+
from phased_array_systems.scenarios.base import ScenarioBase
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
class CommsLinkScenario(ScenarioBase):
|
|
11
|
+
"""Scenario for communications link budget analysis.
|
|
12
|
+
|
|
13
|
+
Defines the parameters needed for a point-to-point or
|
|
14
|
+
satellite communications link budget calculation.
|
|
15
|
+
|
|
16
|
+
Attributes:
|
|
17
|
+
freq_hz: Operating frequency (Hz)
|
|
18
|
+
bandwidth_hz: Signal bandwidth (Hz)
|
|
19
|
+
range_m: Link range/distance (meters)
|
|
20
|
+
required_snr_db: Required SNR for demodulation (dB)
|
|
21
|
+
scan_angle_deg: Beam scan angle from boresight (degrees)
|
|
22
|
+
rx_antenna_gain_db: Receive antenna gain (dB), None for isotropic
|
|
23
|
+
rx_noise_temp_k: Receive system noise temperature (K)
|
|
24
|
+
path_loss_model: Propagation model to use
|
|
25
|
+
atmospheric_loss_db: Additional atmospheric losses (dB)
|
|
26
|
+
rain_loss_db: Rain fade margin (dB)
|
|
27
|
+
polarization_loss_db: Polarization mismatch loss (dB)
|
|
28
|
+
"""
|
|
29
|
+
|
|
30
|
+
bandwidth_hz: float = Field(gt=0, description="Signal bandwidth (Hz)")
|
|
31
|
+
range_m: float = Field(gt=0, description="Link range (m)")
|
|
32
|
+
required_snr_db: float = Field(description="Required SNR (dB)")
|
|
33
|
+
scan_angle_deg: float = Field(
|
|
34
|
+
default=0.0, ge=0, le=90, description="Scan angle from boresight (deg)"
|
|
35
|
+
)
|
|
36
|
+
rx_antenna_gain_db: float | None = Field(
|
|
37
|
+
default=None, description="RX antenna gain (dB), None for isotropic"
|
|
38
|
+
)
|
|
39
|
+
rx_noise_temp_k: float = Field(
|
|
40
|
+
default=290.0, gt=0, description="RX system noise temperature (K)"
|
|
41
|
+
)
|
|
42
|
+
path_loss_model: Literal["fspl"] = Field(
|
|
43
|
+
default="fspl", description="Path loss model"
|
|
44
|
+
)
|
|
45
|
+
atmospheric_loss_db: float = Field(
|
|
46
|
+
default=0.0, ge=0, description="Atmospheric loss (dB)"
|
|
47
|
+
)
|
|
48
|
+
rain_loss_db: float = Field(default=0.0, ge=0, description="Rain loss margin (dB)")
|
|
49
|
+
polarization_loss_db: float = Field(
|
|
50
|
+
default=0.0, ge=0, description="Polarization loss (dB)"
|
|
51
|
+
)
|
|
52
|
+
|
|
53
|
+
@property
|
|
54
|
+
def total_extra_loss_db(self) -> float:
|
|
55
|
+
"""Total additional losses beyond free space path loss."""
|
|
56
|
+
return self.atmospheric_loss_db + self.rain_loss_db + self.polarization_loss_db
|
|
@@ -0,0 +1,42 @@
|
|
|
1
|
+
"""Radar detection scenario definition (stub for Phase 3)."""
|
|
2
|
+
|
|
3
|
+
from typing import Literal
|
|
4
|
+
|
|
5
|
+
from pydantic import Field
|
|
6
|
+
|
|
7
|
+
from phased_array_systems.scenarios.base import ScenarioBase
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
class RadarDetectionScenario(ScenarioBase):
|
|
11
|
+
"""Scenario for radar detection analysis.
|
|
12
|
+
|
|
13
|
+
Note: Full implementation planned for Phase 3.
|
|
14
|
+
|
|
15
|
+
Attributes:
|
|
16
|
+
freq_hz: Operating frequency (Hz)
|
|
17
|
+
bandwidth_hz: Signal bandwidth (Hz)
|
|
18
|
+
range_m: Target range (meters)
|
|
19
|
+
target_rcs_dbsm: Target radar cross section (dBsm)
|
|
20
|
+
pfa: Probability of false alarm
|
|
21
|
+
pd_required: Required probability of detection
|
|
22
|
+
n_pulses: Number of pulses integrated
|
|
23
|
+
scan_angle_deg: Beam scan angle from boresight (degrees)
|
|
24
|
+
integration_type: Coherent or non-coherent integration
|
|
25
|
+
"""
|
|
26
|
+
|
|
27
|
+
bandwidth_hz: float = Field(gt=0, description="Signal bandwidth (Hz)")
|
|
28
|
+
range_m: float = Field(gt=0, description="Target range (m)")
|
|
29
|
+
target_rcs_dbsm: float = Field(description="Target RCS (dBsm)")
|
|
30
|
+
pfa: float = Field(
|
|
31
|
+
default=1e-6, gt=0, lt=1, description="Probability of false alarm"
|
|
32
|
+
)
|
|
33
|
+
pd_required: float = Field(
|
|
34
|
+
default=0.9, gt=0, lt=1, description="Required probability of detection"
|
|
35
|
+
)
|
|
36
|
+
n_pulses: int = Field(default=1, ge=1, description="Number of pulses integrated")
|
|
37
|
+
scan_angle_deg: float = Field(
|
|
38
|
+
default=0.0, ge=0, le=90, description="Scan angle from boresight (deg)"
|
|
39
|
+
)
|
|
40
|
+
integration_type: Literal["coherent", "noncoherent"] = Field(
|
|
41
|
+
default="noncoherent", description="Integration type"
|
|
42
|
+
)
|
|
@@ -0,0 +1,16 @@
|
|
|
1
|
+
"""Trade study tools: DOE, batch evaluation, Pareto analysis."""
|
|
2
|
+
|
|
3
|
+
from phased_array_systems.trades.design_space import DesignSpace, DesignVariable
|
|
4
|
+
from phased_array_systems.trades.doe import generate_doe
|
|
5
|
+
from phased_array_systems.trades.pareto import extract_pareto, filter_feasible, rank_pareto
|
|
6
|
+
from phased_array_systems.trades.runner import BatchRunner
|
|
7
|
+
|
|
8
|
+
__all__ = [
|
|
9
|
+
"DesignSpace",
|
|
10
|
+
"DesignVariable",
|
|
11
|
+
"generate_doe",
|
|
12
|
+
"BatchRunner",
|
|
13
|
+
"extract_pareto",
|
|
14
|
+
"filter_feasible",
|
|
15
|
+
"rank_pareto",
|
|
16
|
+
]
|
|
@@ -0,0 +1,241 @@
|
|
|
1
|
+
"""Design space definition for DOE studies."""
|
|
2
|
+
|
|
3
|
+
from typing import Any, Literal
|
|
4
|
+
|
|
5
|
+
import numpy as np
|
|
6
|
+
import pandas as pd
|
|
7
|
+
from pydantic import BaseModel, Field, model_validator
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
class DesignVariable(BaseModel):
|
|
11
|
+
"""Definition of a single design variable.
|
|
12
|
+
|
|
13
|
+
Supports continuous (float), discrete (int), and categorical variables.
|
|
14
|
+
|
|
15
|
+
Attributes:
|
|
16
|
+
name: Variable name, typically a dot-path like "array.nx"
|
|
17
|
+
type: Variable type ("int", "float", or "categorical")
|
|
18
|
+
low: Lower bound for continuous/discrete variables
|
|
19
|
+
high: Upper bound for continuous/discrete variables
|
|
20
|
+
values: List of allowed values for categorical variables
|
|
21
|
+
"""
|
|
22
|
+
|
|
23
|
+
name: str = Field(description="Variable name (e.g., 'array.nx')")
|
|
24
|
+
type: Literal["int", "float", "categorical"] = "float"
|
|
25
|
+
low: float | None = Field(default=None, description="Lower bound")
|
|
26
|
+
high: float | None = Field(default=None, description="Upper bound")
|
|
27
|
+
values: list[Any] | None = Field(default=None, description="Categorical values")
|
|
28
|
+
|
|
29
|
+
@model_validator(mode="after")
|
|
30
|
+
def validate_bounds_or_values(self) -> "DesignVariable":
|
|
31
|
+
"""Ensure proper bounds/values are set based on type."""
|
|
32
|
+
if self.type in ("int", "float"):
|
|
33
|
+
if self.low is None or self.high is None:
|
|
34
|
+
raise ValueError(f"Variable '{self.name}': low and high required for {self.type}")
|
|
35
|
+
if self.low > self.high:
|
|
36
|
+
raise ValueError(f"Variable '{self.name}': low must be <= high")
|
|
37
|
+
elif self.type == "categorical":
|
|
38
|
+
if not self.values or len(self.values) < 1:
|
|
39
|
+
raise ValueError(f"Variable '{self.name}': values required for categorical")
|
|
40
|
+
return self
|
|
41
|
+
|
|
42
|
+
def sample_uniform(self, n: int, rng: np.random.Generator) -> np.ndarray:
|
|
43
|
+
"""Generate uniform random samples.
|
|
44
|
+
|
|
45
|
+
Args:
|
|
46
|
+
n: Number of samples
|
|
47
|
+
rng: NumPy random generator
|
|
48
|
+
|
|
49
|
+
Returns:
|
|
50
|
+
Array of sampled values
|
|
51
|
+
"""
|
|
52
|
+
if self.type == "float":
|
|
53
|
+
return rng.uniform(self.low, self.high, n)
|
|
54
|
+
elif self.type == "int":
|
|
55
|
+
return rng.integers(int(self.low), int(self.high) + 1, n)
|
|
56
|
+
else: # categorical
|
|
57
|
+
indices = rng.integers(0, len(self.values), n)
|
|
58
|
+
return np.array([self.values[i] for i in indices])
|
|
59
|
+
|
|
60
|
+
def scale_from_unit(self, unit_values: np.ndarray) -> np.ndarray:
|
|
61
|
+
"""Scale values from [0, 1] to actual variable range.
|
|
62
|
+
|
|
63
|
+
Args:
|
|
64
|
+
unit_values: Values in [0, 1]
|
|
65
|
+
|
|
66
|
+
Returns:
|
|
67
|
+
Scaled values in variable's actual range
|
|
68
|
+
"""
|
|
69
|
+
if self.type == "float":
|
|
70
|
+
return self.low + unit_values * (self.high - self.low)
|
|
71
|
+
elif self.type == "int":
|
|
72
|
+
scaled = self.low + unit_values * (self.high - self.low + 1)
|
|
73
|
+
return np.floor(scaled).astype(int).clip(int(self.low), int(self.high))
|
|
74
|
+
else: # categorical
|
|
75
|
+
indices = np.floor(unit_values * len(self.values)).astype(int)
|
|
76
|
+
indices = indices.clip(0, len(self.values) - 1)
|
|
77
|
+
return np.array([self.values[i] for i in indices])
|
|
78
|
+
|
|
79
|
+
def get_grid_values(self, n_levels: int) -> list[Any]:
|
|
80
|
+
"""Get grid values for this variable.
|
|
81
|
+
|
|
82
|
+
Args:
|
|
83
|
+
n_levels: Number of levels for grid
|
|
84
|
+
|
|
85
|
+
Returns:
|
|
86
|
+
List of values at each level
|
|
87
|
+
"""
|
|
88
|
+
if self.type == "float":
|
|
89
|
+
return list(np.linspace(self.low, self.high, n_levels))
|
|
90
|
+
elif self.type == "int":
|
|
91
|
+
# For integers, use actual integer values
|
|
92
|
+
all_ints = list(range(int(self.low), int(self.high) + 1))
|
|
93
|
+
if len(all_ints) <= n_levels:
|
|
94
|
+
return all_ints
|
|
95
|
+
# Subsample evenly
|
|
96
|
+
indices = np.linspace(0, len(all_ints) - 1, n_levels).astype(int)
|
|
97
|
+
return [all_ints[i] for i in indices]
|
|
98
|
+
else: # categorical
|
|
99
|
+
return list(self.values)
|
|
100
|
+
|
|
101
|
+
|
|
102
|
+
class DesignSpace(BaseModel):
|
|
103
|
+
"""Collection of design variables defining a design space.
|
|
104
|
+
|
|
105
|
+
Provides methods for sampling the design space using various
|
|
106
|
+
DOE methods (grid, random, LHS).
|
|
107
|
+
|
|
108
|
+
Attributes:
|
|
109
|
+
variables: List of design variables
|
|
110
|
+
name: Optional name for the design space
|
|
111
|
+
"""
|
|
112
|
+
|
|
113
|
+
variables: list[DesignVariable] = Field(default_factory=list)
|
|
114
|
+
name: str | None = None
|
|
115
|
+
|
|
116
|
+
def add_variable(
|
|
117
|
+
self,
|
|
118
|
+
name: str,
|
|
119
|
+
type: Literal["int", "float", "categorical"] = "float",
|
|
120
|
+
low: float | None = None,
|
|
121
|
+
high: float | None = None,
|
|
122
|
+
values: list[Any] | None = None,
|
|
123
|
+
) -> "DesignSpace":
|
|
124
|
+
"""Add a variable to the design space (fluent interface).
|
|
125
|
+
|
|
126
|
+
Args:
|
|
127
|
+
name: Variable name
|
|
128
|
+
type: Variable type
|
|
129
|
+
low: Lower bound
|
|
130
|
+
high: Upper bound
|
|
131
|
+
values: Categorical values
|
|
132
|
+
|
|
133
|
+
Returns:
|
|
134
|
+
Self for chaining
|
|
135
|
+
"""
|
|
136
|
+
var = DesignVariable(name=name, type=type, low=low, high=high, values=values)
|
|
137
|
+
self.variables.append(var)
|
|
138
|
+
return self
|
|
139
|
+
|
|
140
|
+
@property
|
|
141
|
+
def n_dims(self) -> int:
|
|
142
|
+
"""Number of dimensions (variables) in the design space."""
|
|
143
|
+
return len(self.variables)
|
|
144
|
+
|
|
145
|
+
@property
|
|
146
|
+
def variable_names(self) -> list[str]:
|
|
147
|
+
"""List of variable names."""
|
|
148
|
+
return [v.name for v in self.variables]
|
|
149
|
+
|
|
150
|
+
def sample(
|
|
151
|
+
self,
|
|
152
|
+
method: Literal["grid", "random", "lhs"] = "lhs",
|
|
153
|
+
n_samples: int = 100,
|
|
154
|
+
seed: int | None = None,
|
|
155
|
+
grid_levels: int | list[int] | None = None,
|
|
156
|
+
) -> pd.DataFrame:
|
|
157
|
+
"""Sample the design space.
|
|
158
|
+
|
|
159
|
+
Args:
|
|
160
|
+
method: Sampling method ("grid", "random", "lhs")
|
|
161
|
+
n_samples: Number of samples (ignored for grid method)
|
|
162
|
+
seed: Random seed for reproducibility
|
|
163
|
+
grid_levels: Number of levels per variable for grid method
|
|
164
|
+
|
|
165
|
+
Returns:
|
|
166
|
+
DataFrame with columns for each variable plus 'case_id'
|
|
167
|
+
"""
|
|
168
|
+
if method == "grid":
|
|
169
|
+
return self._sample_grid(grid_levels)
|
|
170
|
+
elif method == "random":
|
|
171
|
+
return self._sample_random(n_samples, seed)
|
|
172
|
+
elif method == "lhs":
|
|
173
|
+
return self._sample_lhs(n_samples, seed)
|
|
174
|
+
else:
|
|
175
|
+
raise ValueError(f"Unknown sampling method: {method}")
|
|
176
|
+
|
|
177
|
+
def _sample_grid(self, grid_levels: int | list[int] | None) -> pd.DataFrame:
|
|
178
|
+
"""Generate full factorial grid."""
|
|
179
|
+
if grid_levels is None:
|
|
180
|
+
grid_levels = 3 # Default
|
|
181
|
+
|
|
182
|
+
if isinstance(grid_levels, int):
|
|
183
|
+
levels_per_var = [grid_levels] * len(self.variables)
|
|
184
|
+
else:
|
|
185
|
+
levels_per_var = grid_levels
|
|
186
|
+
|
|
187
|
+
# Generate grid values for each variable
|
|
188
|
+
var_values = []
|
|
189
|
+
for var, n_levels in zip(self.variables, levels_per_var, strict=True):
|
|
190
|
+
var_values.append(var.get_grid_values(n_levels))
|
|
191
|
+
|
|
192
|
+
# Create full factorial grid using meshgrid
|
|
193
|
+
grids = np.meshgrid(*var_values, indexing="ij")
|
|
194
|
+
flat_grids = [g.flatten() for g in grids]
|
|
195
|
+
|
|
196
|
+
# Build DataFrame
|
|
197
|
+
data = {var.name: flat_grids[i] for i, var in enumerate(self.variables)}
|
|
198
|
+
df = pd.DataFrame(data)
|
|
199
|
+
|
|
200
|
+
# Add case IDs
|
|
201
|
+
df.insert(0, "case_id", [f"case_{i:05d}" for i in range(len(df))])
|
|
202
|
+
|
|
203
|
+
return df
|
|
204
|
+
|
|
205
|
+
def _sample_random(self, n_samples: int, seed: int | None) -> pd.DataFrame:
|
|
206
|
+
"""Generate random samples."""
|
|
207
|
+
rng = np.random.default_rng(seed)
|
|
208
|
+
|
|
209
|
+
data = {}
|
|
210
|
+
for var in self.variables:
|
|
211
|
+
data[var.name] = var.sample_uniform(n_samples, rng)
|
|
212
|
+
|
|
213
|
+
df = pd.DataFrame(data)
|
|
214
|
+
df.insert(0, "case_id", [f"case_{i:05d}" for i in range(len(df))])
|
|
215
|
+
|
|
216
|
+
return df
|
|
217
|
+
|
|
218
|
+
def _sample_lhs(self, n_samples: int, seed: int | None) -> pd.DataFrame:
|
|
219
|
+
"""Generate Latin Hypercube samples."""
|
|
220
|
+
from scipy.stats import qmc
|
|
221
|
+
|
|
222
|
+
# Generate LHS in unit hypercube
|
|
223
|
+
sampler = qmc.LatinHypercube(d=len(self.variables), seed=seed)
|
|
224
|
+
unit_samples = sampler.random(n_samples)
|
|
225
|
+
|
|
226
|
+
# Scale to actual variable ranges
|
|
227
|
+
data = {}
|
|
228
|
+
for i, var in enumerate(self.variables):
|
|
229
|
+
data[var.name] = var.scale_from_unit(unit_samples[:, i])
|
|
230
|
+
|
|
231
|
+
df = pd.DataFrame(data)
|
|
232
|
+
df.insert(0, "case_id", [f"case_{i:05d}" for i in range(len(df))])
|
|
233
|
+
|
|
234
|
+
return df
|
|
235
|
+
|
|
236
|
+
def get_variable(self, name: str) -> DesignVariable | None:
|
|
237
|
+
"""Get a variable by name."""
|
|
238
|
+
for var in self.variables:
|
|
239
|
+
if var.name == name:
|
|
240
|
+
return var
|
|
241
|
+
return None
|