Require function as first argument to cram.deplete

Done in order to remove hardcoded used of CRAM48
This commit is contained in:
Andrew Johnson 2020-05-09 10:26:14 -04:00
parent 2f0a894131
commit 24212ae731
No known key found for this signature in database
GPG key ID: 253418E91B7F6FEB
3 changed files with 37 additions and 30 deletions

View file

@ -20,11 +20,14 @@ __all__ = [
"Cram16Solver", "Cram48Solver", "IPFCramSolver"]
def deplete(chain, x, rates, dt, matrix_func=None):
def deplete(func, chain, x, rates, dt, matrix_func=None):
"""Deplete materials using given reaction rates for a specified time
Parameters
----------
func : callable
Function to use to get new compositions. Expected to have the
signature ``func(A, n0, t) -> n1``
chain : openmc.deplete.Chain
Depletion chain
x : list of numpy.ndarray
@ -63,7 +66,7 @@ def deplete(chain, x, rates, dt, matrix_func=None):
# Use multiprocessing pool to distribute work
with Pool() as pool:
inputs = zip(matrices, x, repeat(dt))
x_result = list(pool.starmap(CRAM48, inputs))
x_result = list(pool.starmap(func, inputs))
return x_result

View file

@ -2,7 +2,7 @@ import copy
from itertools import repeat
from .abc import Integrator, SIIntegrator, OperatorResult
from .cram import timed_deplete
from .cram import CRAM48, timed_deplete
from ._matrix_funcs import (
cf4_f1, cf4_f2, cf4_f3, cf4_f4, celi_f1, celi_f2,
leqi_f1, leqi_f2, leqi_f3, leqi_f4, rk4_f1, rk4_f4
@ -96,7 +96,8 @@ class PredictorIntegrator(Integrator):
operator with predictor
"""
proc_time, conc_end = timed_deplete(self.chain, conc, rates, dt)
proc_time, conc_end = timed_deplete(
CRAM48, self.chain, conc, rates, dt)
return proc_time, [conc_end], []
@ -187,12 +188,14 @@ class CECMIntegrator(Integrator):
Eigenvalue and reaction rates from transport simulations
"""
# deplete across first half of inteval
time0, x_middle = timed_deplete(self.chain, conc, rates, dt / 2)
time0, x_middle = timed_deplete(
CRAM48, self.chain, conc, rates, dt / 2)
res_middle = self.operator(x_middle, power)
# deplete across entire interval with BOS concentrations,
# MOS reaction rates
time1, x_end = timed_deplete(self.chain, conc, res_middle.rates, dt)
time1, x_end = timed_deplete(
CRAM48, self.chain, conc, res_middle.rates, dt)
return time0 + time1, [x_middle, x_end], [res_middle]
@ -288,26 +291,26 @@ class CF4Integrator(Integrator):
"""
# Step 1: deplete with matrix 1/2*A(y0)
time1, conc_eos1 = timed_deplete(
self.chain, bos_conc, bos_rates, dt, matrix_func=cf4_f1)
CRAM48, self.chain, bos_conc, bos_rates, dt, matrix_func=cf4_f1)
res1 = self.operator(conc_eos1, power)
# Step 2: deplete with matrix 1/2*A(y1)
time2, conc_eos2 = timed_deplete(
self.chain, bos_conc, res1.rates, dt, matrix_func=cf4_f1)
CRAM48, self.chain, bos_conc, res1.rates, dt, matrix_func=cf4_f1)
res2 = self.operator(conc_eos2, power)
# Step 3: deplete with matrix -1/2*A(y0)+A(y2)
list_rates = list(zip(bos_rates, res2.rates))
time3, conc_eos3 = timed_deplete(
self.chain, conc_eos1, list_rates, dt, matrix_func=cf4_f2)
CRAM48, self.chain, conc_eos1, list_rates, dt, matrix_func=cf4_f2)
res3 = self.operator(conc_eos3, power)
# Step 4: deplete with two matrix exponentials
list_rates = list(zip(bos_rates, res1.rates, res2.rates, res3.rates))
time4, conc_inter = timed_deplete(
self.chain, bos_conc, list_rates, dt, matrix_func=cf4_f3)
CRAM48, self.chain, bos_conc, list_rates, dt, matrix_func=cf4_f3)
time5, conc_eos5 = timed_deplete(
self.chain, conc_inter, list_rates, dt, matrix_func=cf4_f4)
CRAM48, self.chain, conc_inter, list_rates, dt, matrix_func=cf4_f4)
return (time1 + time2 + time3 + time4 + time5,
[conc_eos1, conc_eos2, conc_eos3, conc_eos5],
@ -403,17 +406,18 @@ class CELIIntegrator(Integrator):
simulation
"""
# deplete to end using BOS rates
proc_time, conc_ce = timed_deplete(self.chain, bos_conc, rates, dt)
proc_time, conc_ce = timed_deplete(
CRAM48, self.chain, bos_conc, rates, dt)
res_ce = self.operator(conc_ce, power)
# deplete using two matrix exponentials
list_rates = list(zip(rates, res_ce.rates))
time_le1, conc_inter = timed_deplete(
self.chain, bos_conc, list_rates, dt, matrix_func=celi_f1)
CRAM48, self.chain, bos_conc, list_rates, dt, matrix_func=celi_f1)
time_le2, conc_end = timed_deplete(
self.chain, conc_inter, list_rates, dt, matrix_func=celi_f2)
CRAM48, self.chain, conc_inter, list_rates, dt, matrix_func=celi_f2)
return proc_time + time_le1 + time_le1, [conc_ce, conc_end], [res_ce]
@ -508,23 +512,23 @@ class EPCRK4Integrator(Integrator):
# Step 1: deplete with matrix A(y0) / 2
time1, conc1 = timed_deplete(
self.chain, conc, rates, dt, matrix_func=rk4_f1)
CRAM48, self.chain, conc, rates, dt, matrix_func=rk4_f1)
res1 = self.operator(conc1, power)
# Step 2: deplete with matrix A(y1) / 2
time2, conc2 = timed_deplete(
self.chain, conc, res1.rates, dt, matrix_func=rk4_f1)
CRAM48, self.chain, conc, res1.rates, dt, matrix_func=rk4_f1)
res2 = self.operator(conc2, power)
# Step 3: deplete with matrix A(y2)
time3, conc3 = timed_deplete(
self.chain, conc, res2.rates, dt)
CRAM48, self.chain, conc, res2.rates, dt)
res3 = self.operator(conc3, power)
# Step 4: deplete with matrix built from weighted rates
list_rates = list(zip(rates, res1.rates, res2.rates, res3.rates))
time4, conc4 = timed_deplete(
self.chain, conc, list_rates, dt, matrix_func=rk4_f4)
CRAM48, self.chain, conc, list_rates, dt, matrix_func=rk4_f4)
return (time1 + time2 + time3 + time4, [conc1, conc2, conc3, conc4],
[res1, res2, res3])
@ -646,9 +650,9 @@ class LEQIIntegrator(Integrator):
self._prev_rates, bos_res.rates, repeat(prev_dt), repeat(dt)))
time1, conc_inter = timed_deplete(
self.chain, bos_conc, le_inputs, dt, matrix_func=leqi_f1)
CRAM48, self.chain, bos_conc, le_inputs, dt, matrix_func=leqi_f1)
time2, conc_eos0 = timed_deplete(
self.chain, conc_inter, le_inputs, dt, matrix_func=leqi_f2)
CRAM48, self.chain, conc_inter, le_inputs, dt, matrix_func=leqi_f2)
res_inter = self.operator(conc_eos0, power)
@ -657,9 +661,9 @@ class LEQIIntegrator(Integrator):
repeat(prev_dt), repeat(dt)))
time3, conc_inter = timed_deplete(
self.chain, bos_conc, qi_inputs, dt, matrix_func=leqi_f3)
CRAM48, self.chain, bos_conc, qi_inputs, dt, matrix_func=leqi_f3)
time4, conc_eos1 = timed_deplete(
self.chain, conc_inter, qi_inputs, dt, matrix_func=leqi_f4)
CRAM48, self.chain, conc_inter, qi_inputs, dt, matrix_func=leqi_f4)
# store updated rates
self._prev_rates = copy.deepcopy(bos_res.rates)
@ -754,7 +758,7 @@ class SICELIIntegrator(SIIntegrator):
simulations
"""
proc_time, eos_conc = timed_deplete(
self.chain, bos_conc, bos_rates, dt)
CRAM48, self.chain, bos_conc, bos_rates, dt)
inter_conc = copy.deepcopy(eos_conc)
# Begin iteration
@ -770,9 +774,9 @@ class SICELIIntegrator(SIIntegrator):
list_rates = list(zip(bos_rates, res_bar.rates))
time1, inter_conc = timed_deplete(
self.chain, bos_conc, list_rates, dt, matrix_func=celi_f1)
CRAM48, self.chain, bos_conc, list_rates, dt, matrix_func=celi_f1)
time2, inter_conc = timed_deplete(
self.chain, inter_conc, list_rates, dt, matrix_func=celi_f2)
CRAM48, self.chain, inter_conc, list_rates, dt, matrix_func=celi_f2)
proc_time += time1 + time2
# end iteration
@ -880,9 +884,9 @@ class SILEQIIntegrator(SIIntegrator):
inputs = list(zip(self._prev_rates, bos_rates,
repeat(prev_dt), repeat(dt)))
proc_time, inter_conc = timed_deplete(
self.chain, bos_conc, inputs, dt, matrix_func=leqi_f1)
CRAM48, self.chain, bos_conc, inputs, dt, matrix_func=leqi_f1)
time1, eos_conc = timed_deplete(
self.chain, inter_conc, inputs, dt, matrix_func=leqi_f2)
CRAM48, self.chain, inter_conc, inputs, dt, matrix_func=leqi_f2)
proc_time += time1
inter_conc = copy.deepcopy(eos_conc)
@ -900,9 +904,9 @@ class SILEQIIntegrator(SIIntegrator):
inputs = list(zip(self._prev_rates, bos_rates, res_bar.rates,
repeat(prev_dt), repeat(dt)))
time1, inter_conc = timed_deplete(
self.chain, bos_conc, inputs, dt, matrix_func=leqi_f3)
CRAM48, self.chain, bos_conc, inputs, dt, matrix_func=leqi_f3)
time2, inter_conc = timed_deplete(
self.chain, inter_conc, inputs, dt, matrix_func=leqi_f4)
CRAM48, self.chain, inter_conc, inputs, dt, matrix_func=leqi_f4)
proc_time += time1 + time2
return proc_time, [eos_conc, inter_conc], [res_bar]

View file

@ -414,7 +414,7 @@ def test_fission_yield_attribute(simple_chain):
dummy_conc = [[1, 2]] * (len(empty_chain.fission_yields) + 1)
with pytest.raises(
ValueError, match="fission yield.*not equal.*compositions"):
cram.deplete(empty_chain, dummy_conc, None, 0.5)
cram.deplete(cram.CRAM48, empty_chain, dummy_conc, None, 0.5)
def test_validate(simple_chain):