Skip to content

Commit c2c3c40

Browse files
committed
enable new QRS functions
1 parent 362e4f5 commit c2c3c40

6 files changed

Lines changed: 318 additions & 20 deletions

File tree

pyxtal/interface/ase_opt.py

Lines changed: 17 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -370,7 +370,19 @@ def handler(signum, frame):
370370

371371
try:
372372
atoms.calc = get_calculator(calculator, model=model, quick=quick)
373-
atoms.set_constraint(FixSymmetry(atoms))
373+
try:
374+
atoms.set_constraint(FixSymmetry(atoms))
375+
except (TypeError, AttributeError, RuntimeError) as exc:
376+
# spglib can return None for some lattices; continue without symmetry constraint
377+
try:
378+
logger.warning(
379+
f"Warning {label} FixSymmetry failed ({exc}); relaxing without symmetry constraint"
380+
)
381+
except Exception:
382+
print(
383+
f"Warning {label} FixSymmetry failed ({exc}); relaxing without symmetry constraint",
384+
flush=True,
385+
)
374386

375387
if opt_lat:
376388
ecf = UnitCellFilter(atoms)
@@ -394,11 +406,10 @@ def handler(signum, frame):
394406
atoms = None
395407

396408
except TimeoutError:
397-
logger.warning(f"Warning {label} timed out after {timeout} seconds.")
398-
atoms = None
399-
400-
except TypeError:
401-
logger.warning(f"Warning {label} spglib error in getting the lattice")
409+
try:
410+
logger.warning(f"Warning {label} timed out after {timeout} seconds.")
411+
except Exception:
412+
print(f"Warning {label} timed out after {timeout} seconds.", flush=True)
402413
atoms = None
403414

404415
finally:

pyxtal/optimize/QRS.py

Lines changed: 89 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -262,6 +262,7 @@ def expand_delta_angle_per_dof(delta_spec, composition, wp_bounds, default=30.0)
262262
"""Expand per-component delta specs to a flat per-DOF angle resolution list.
263263
264264
Each component entry may be:
265+
- ``None`` (monoatomic / no orientation DOFs; frac slots get ``default``)
265266
- a scalar (same resolution for all angle DOFs at that site)
266267
- a length-3 sequence ``[alpha, beta, gamma]`` for Euler angles
267268
"""
@@ -280,10 +281,15 @@ def expand_delta_angle_per_dof(delta_spec, composition, wp_bounds, default=30.0)
280281
for _ in range(int(count)):
281282
bounds = wp_bounds[site_idx]
282283
euler = None
283-
if isinstance(spec, (list, tuple)) and len(spec) == 3:
284+
if spec is None:
285+
site_scalar = None
286+
elif isinstance(spec, (list, tuple)) and len(spec) == 3:
284287
euler = [float(d) for d in spec]
288+
site_scalar = None
285289
elif isinstance(spec, (list, tuple)) and len(spec) == 1:
286-
spec = float(spec[0])
290+
site_scalar = float(spec[0])
291+
else:
292+
site_scalar = float(spec)
287293

288294
euler_idx = 0
289295
for lb, ub in bounds:
@@ -294,8 +300,11 @@ def expand_delta_angle_per_dof(delta_spec, composition, wp_bounds, default=30.0)
294300
da = euler[euler_idx] if euler_idx < len(euler) else euler[-1]
295301
flat.append(float(da))
296302
euler_idx += 1
303+
elif site_scalar is None:
304+
# No orientation expected for this component; keep placeholder.
305+
flat.append(float(default))
297306
else:
298-
flat.append(float(spec))
307+
flat.append(float(site_scalar))
299308
site_idx += 1
300309

301310
return flat
@@ -305,10 +314,14 @@ def _delta_angle_enabled(delta_angle, delta_length):
305314
"""Return True when uneven-grid angle resolution should be used."""
306315
if delta_length > 0:
307316
return True
317+
if delta_angle is None:
318+
return False
308319
if isinstance(delta_angle, (int, float)):
309320
return delta_angle > 0
310321
if isinstance(delta_angle, (list, tuple)):
311322
def _positive(spec):
323+
if spec is None:
324+
return False
312325
if isinstance(spec, (list, tuple)):
313326
return any(_positive(x) for x in spec)
314327
return float(spec) > 0
@@ -340,12 +353,14 @@ def compute_wp_resolutions(wp_bounds, cell_lengths, delta_length=1.0, delta_angl
340353
per_dof_delta = None
341354
if isinstance(delta_angle, (list, tuple, np.ndarray)):
342355
if len(delta_angle) == total_dofs:
343-
per_dof_delta = [float(d) for d in delta_angle]
356+
per_dof_delta = [None if d is None else float(d) for d in delta_angle]
344357
elif len(delta_angle) == len(wp_bounds):
345-
per_site_delta = [float(d) for d in delta_angle]
358+
per_site_delta = [None if d is None else float(d) for d in delta_angle]
346359
else:
347360
try:
348-
per_site_delta = [float(delta_angle[0])] * len(wp_bounds)
361+
first = delta_angle[0]
362+
fill = None if first is None else float(first)
363+
per_site_delta = [fill] * len(wp_bounds)
349364
except Exception:
350365
per_site_delta = [float(delta_angle)] * len(wp_bounds)
351366
else:
@@ -364,12 +379,59 @@ def compute_wp_resolutions(wp_bounds, cell_lengths, delta_length=1.0, delta_angl
364379
else: # angle DOF (Euler or torsion)
365380
if per_dof_delta is not None:
366381
site_delta = per_dof_delta[dof_idx]
367-
n = max(1, int(round(span / site_delta)))
382+
if site_delta is None or float(site_delta) <= 0:
383+
n = 1
384+
else:
385+
n = max(1, int(round(span / float(site_delta))))
368386
dof_idx += 1
369387
n_levels.append(n)
370388
return n_levels
371389

372390

391+
def budget_grid_levels(n_levels, max_product, min_level=1):
392+
"""
393+
Reduce per-DOF grid levels so that prod(n_levels) <= max_product.
394+
395+
Uses a uniform log-space scale first (preserves relative fineness), then
396+
decrements the currently largest remaining levels until the product fits.
397+
Levels never drop below ``min_level``.
398+
399+
Args:
400+
n_levels: per-DOF bin counts
401+
max_product: maximum allowed product (None / <=0 disables budgeting)
402+
min_level: minimum bins per DOF (default 1)
403+
404+
Returns:
405+
list[int]: budgeted levels (same length as input)
406+
"""
407+
import math
408+
409+
levels = [max(int(min_level), int(n)) for n in n_levels]
410+
if not levels or max_product is None or max_product <= 0:
411+
return levels
412+
413+
def _log_prod(lv):
414+
return sum(math.log(x) for x in lv)
415+
416+
log_budget = math.log(float(max_product))
417+
if _log_prod(levels) <= log_budget + 1e-12:
418+
return levels
419+
420+
# Uniform scale in log space so relative resolution is roughly preserved.
421+
scale = math.exp((log_budget - _log_prod(levels)) / len(levels))
422+
levels = [max(int(min_level), int(round(n * scale))) for n in levels]
423+
424+
# Trim residual overshoot from the largest dimensions first.
425+
while _log_prod(levels) > log_budget + 1e-12:
426+
candidates = [i for i, n in enumerate(levels) if n > min_level]
427+
if not candidates:
428+
break
429+
i_max = max(candidates, key=lambda j: levels[j])
430+
levels[i_max] -= 1
431+
432+
return levels
433+
434+
373435
def trim_wp_bounds_for_molecules(wp_bounds, composition, molecules):
374436
"""Remove torsion DOFs from wp bounds when fixed conformer pools are provided."""
375437
if molecules is None:
@@ -474,6 +536,8 @@ class QRS(GlobalOptimize):
474536
use_mpi (bool): whether or not use mpi for parallel calculation
475537
delta_length (float): grid resolution in Å for fractional-coordinate DOF (0 = use Sobol)
476538
delta_angle (float): grid resolution in degrees for angle DOF (0 = use Sobol)
539+
max_grid_product (float or None): cap on prod(n_levels); levels are scaled down
540+
if the raw product exceeds this (default: 1e9; None disables)
477541
"""
478542

479543
def __init__(
@@ -515,6 +579,7 @@ def __init__(
515579
delta_length: float = 1.0,
516580
delta_angle: float = 60.0,
517581
max_grid_attempts: int | None = None,
582+
max_grid_product: float | None = 1e9,
518583
):
519584

520585
# POPULATION parameters:
@@ -525,6 +590,7 @@ def __init__(
525590
self.delta_length = delta_length
526591
self.delta_angle = delta_angle
527592
self.max_grid_attempts = max_grid_attempts
593+
self.max_grid_product = max_grid_product
528594

529595
# initialize other base parameters
530596
GlobalOptimize.__init__(
@@ -575,6 +641,8 @@ def full_str(self):
575641
s += f"\nPopulation: {self.N_pop:4d}"
576642
if self.max_grid_attempts is not None:
577643
s += f"\nGrid cap : {self.max_grid_attempts:4d}"
644+
if self.max_grid_product is not None and self.max_grid_product > 0:
645+
s += f"\nGrid budget: {self.max_grid_product:.3g}"
578646
return s
579647

580648
def _init_qrs_params(self):
@@ -626,7 +694,9 @@ def _init_qrs_params(self):
626694
else:
627695
per_site_da = []
628696
for comp_idx, cnt in enumerate(self.composition):
629-
per_site_da.extend([da[comp_idx]] * int(cnt))
697+
# None = monoatomic (no orientation); unused for angle bins.
698+
val = da[comp_idx]
699+
per_site_da.extend([val] * int(cnt))
630700
per_dof_da = None
631701
else:
632702
per_site_da = da
@@ -638,7 +708,17 @@ def _init_qrs_params(self):
638708
dl,
639709
per_dof_da if per_dof_da is not None else per_site_da,
640710
)
641-
print(f"Computed per-DOF grid levels: {n_levels}")
711+
raw_product = int(np.prod(n_levels)) if n_levels else 0
712+
print(f"Computed per-DOF grid levels: {n_levels} (product={raw_product})")
713+
if self.max_grid_product is not None and self.max_grid_product > 0:
714+
budgeted = budget_grid_levels(n_levels, self.max_grid_product)
715+
if budgeted != list(n_levels):
716+
budgeted_product = int(np.prod(budgeted)) if budgeted else 0
717+
print(
718+
f"Budgeted grid levels to product <= {self.max_grid_product:.3g}: "
719+
f"{budgeted} (product={budgeted_product})"
720+
)
721+
n_levels = budgeted
642722
self.sampler = GridSampler(n_levels)
643723
print(f"GridSampler initialised: {self.sampler}")
644724
else:

pyxtal/optimize/base.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -385,7 +385,7 @@ def __str__(self):
385385
s += f"\nN_torsion : {self.N_torsion:d}"
386386
s += f"\nsg : {self.sg!s:s}"
387387
s += f"\nncpu : {self.size:d}"
388-
s += f"\ndiretory : {self.workdir:s}"
388+
s += f"\ndirectory : {self.workdir:s}"
389389
s += f"\nopt_lat : {self.opt_lat!s:s}"
390390
s += f"\nusp_mpi : {self.use_mpi!s:s}\n"
391391
s += f"\nmlp : {self.mlp!s:s}\n"

pyxtal/optimize/common.py

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -50,8 +50,10 @@ def _resolve_delta_angle(delta_angle, default=15.0):
5050
if isinstance(delta_angle, (list, tuple, np.ndarray)):
5151
flat = []
5252
for item in delta_angle:
53+
if item is None:
54+
continue
5355
if isinstance(item, (list, tuple, np.ndarray)):
54-
flat.extend(float(x) for x in item)
56+
flat.extend(float(x) for x in item if x is not None)
5557
else:
5658
flat.append(float(item))
5759
return min(flat) if flat else default

pyxtal/wyckoff_site.py

Lines changed: 8 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -721,11 +721,16 @@ def __repr__(self):
721721
def _get_dof(self):
722722
"""
723723
get the number of dof for the given wyckoff site:
724+
725+
Fractional Wyckoff freedom plus orientation/torsions for multi-atom
726+
molecules. Monoatomic species have no orientation DOFs.
724727
"""
725728
dof = np.linalg.matrix_rank(self.wp.ops[0].rotation_matrix)
726-
self.dof = dof + 3
727-
if self.molecule.torsionlist is not None:
728-
self.dof += len(self.molecule.torsionlist)
729+
if len(self.molecule.mol) > 1:
730+
dof += 3
731+
if self.molecule.torsionlist is not None:
732+
dof += len(self.molecule.torsionlist)
733+
self.dof = dof
729734

730735
def get_bounds(self):
731736
"""

0 commit comments

Comments
 (0)