Generated from the full canonical file for this source snapshot. Line numbers match the library source.
Source SHA256: 5ea6314bb2be68b06e574b03ba37f1608c237bf02332b32ae6290d3d053ecf26
1"""Detector calibration, zero-extended spatial response and separate observations.23No estimator differentiates a realised count draw. Poisson draws use counter4addresses shared with transport, a bounded rejection loop and an explicit error5on budget exhaustion. There is no Gaussian replacement for a Poisson tail.6"""78# Warp annotations are executable DSL expressions; host interfaces remain strict.9# The optional GPU import is resolved only when an operator is prepared.10# pyright: reportInvalidTypeForm=false, reportUnknownParameterType=false11# pyright: reportUnknownMemberType=false, reportUnknownArgumentType=false12# pyright: reportUnknownVariableType=false, reportUntypedFunctionDecorator=false13# pyright: reportMissingImports=false, reportUntypedClassDecorator=false1415from functools import cache1617import warp as wp1819from dpt.kernels.random import random420from dpt.kernels.spectral import checked_store2122STRICT = {"fast_math": False, "fuse_fp": True, "enable_backward": False}23wp.set_module_options(STRICT)24TILE = 256252627@wp.kernel28def check_positive(values: wp.array(dtype=wp.float32), status: wp.array(dtype=wp.int32)):29p = wp.tid()30if not wp.isfinite(values[p]) or values[p] <= wp.float32(0.0):31wp.atomic_or(status, 0, 1)323334@wp.kernel35def check_rates(values: wp.array(dtype=wp.float32), status: wp.array(dtype=wp.int32)):36p = wp.tid()37if not wp.isfinite(values[p]) or values[p] < wp.float32(0.0) or values[p] > wp.float32(1.0e9):38wp.atomic_or(status, 0, 1)394041# region book:detector-calibration42@cache43def get_calibration(shared_gain: bool, shared_offset: bool):44@wp.kernel(module="unique", module_options=STRICT)45def calibration(46mean: wp.array(dtype=wp.float32),47gain: wp.array(dtype=wp.float32),48exposure: wp.array(dtype=wp.float32),49offset: wp.array(dtype=wp.float32),50output: wp.array(dtype=wp.float32),51status: wp.array(dtype=wp.int32),52):53p = wp.tid()54gi = p55oi = p56if wp.static(shared_gain):57gi = 058if wp.static(shared_offset):59oi = 060signal = wp.float64(gain[gi]) * wp.float64(exposure[0]) * wp.float64(mean[p]) + wp.float64(61offset[oi]62)63output[p] = checked_store(signal, status)6465return calibration666768# endregion book:detector-calibration697071@cache72def get_calibration_pixel_vjp(shared_gain: bool, write_mean: bool):73@wp.kernel(module="unique", module_options=STRICT)74def vjp(75gain: wp.array(dtype=wp.float32),76exposure: wp.array(dtype=wp.float32),77seed: wp.array(dtype=wp.float32),78output: wp.array(dtype=wp.float32),79status: wp.array(dtype=wp.int32),80):81p = wp.tid()82gi = p83if wp.static(shared_gain):84gi = 085if wp.static(write_mean):86output[p] = checked_store(87wp.float64(seed[p]) * wp.float64(gain[gi]) * wp.float64(exposure[0]),88status,89)9091return vjp929394@cache95def get_calibration_partials(shared_gain: bool, kind: int):96@wp.kernel(module="unique", module_options=STRICT)97def partials(98mean: wp.array(dtype=wp.float32),99gain: wp.array(dtype=wp.float32),100exposure: wp.array(dtype=wp.float32),101seed: wp.array(dtype=wp.float32),102pixels: int,103groups: int,104output: wp.array(dtype=wp.float64),105):106group, lane = wp.tid()107total = wp.float64(0.0)108compensation = wp.float64(0.0)109p = group * TILE + lane110while p < pixels:111gi = p112if wp.static(shared_gain):113gi = 0114term = wp.float64(seed[p])115if wp.static(kind == 0):116term = term * wp.float64(exposure[0]) * wp.float64(mean[p])117elif wp.static(kind == 1):118term = term * wp.float64(gain[gi]) * wp.float64(mean[p])119corrected = term - compensation120updated = total + corrected121compensation = (updated - total) - corrected122total = updated123# Do not overflow the final int32 stride at maximum image capacity.124if pixels - p <= groups * TILE:125break126p = p + groups * TILE127values = wp.tile(total)128total_tile = wp.tile_sum(values)129wp.tile_store(output, total_tile, offset=group)130131return partials132133134# region book:detector-spatial-transpose135@cache136def get_blur(height: int, width: int, kernel_height: int, kernel_width: int, transpose: bool):137@wp.kernel(module="unique", module_options=STRICT)138def blur(139source: wp.array(dtype=wp.float32),140weights: wp.array(dtype=wp.float64),141output: wp.array(dtype=wp.float32),142status: wp.array(dtype=wp.int32),143):144p = wp.tid()145row = p // width146column = p % width147total = wp.float64(0.0)148compensation = wp.float64(0.0)149for kr in range(kernel_height):150for kc in range(kernel_width):151dr = kr - kernel_height // 2152dc = kc - kernel_width // 2153if wp.static(transpose):154dr = -dr155dc = -dc156sr = row + dr157sc = column + dc158# Zero extension loses signal crossing the finite detector edge.159# Do not renormalise a boundary row: that changes both B and Bᵀ.160if sr >= 0 and sr < height and sc >= 0 and sc < width:161term = wp.float64(weights[kr * kernel_width + kc]) * wp.float64(162source[sr * width + sc]163)164corrected = term - compensation165updated = total + corrected166compensation = (updated - total) - corrected167total = updated168output[p] = checked_store(total, status)169170return blur171172173# endregion book:detector-spatial-transpose174175176@wp.func177def uniform_pair(seed: wp.uint64, identity: wp.uint64, event: wp.uint32, domain: wp.uint32):178"""Two open uniforms using 53 counter bits each, without endpoint clamping."""179draw = random4(seed, identity, event, domain)180scale32 = wp.float64(4294967296.0)181scale21 = wp.float64(2097152.0)182denominator = wp.float64(9007199254740994.0)183first = (184wp.floor(draw[0] * scale32) * scale21 + wp.floor(draw[1] * scale21) + wp.float64(1.0)185) / denominator186second = (187wp.floor(draw[2] * scale32) * scale21 + wp.floor(draw[3] * scale21) + wp.float64(1.0)188) / denominator189return wp.vec2d(first, second)190191192@wp.func_native("return lgamma(value);")193def log_gamma(value: wp.float64) -> wp.float64:194"""Native double log-gamma for the Poisson rejection acceptance inequality."""195...196197198@wp.func199def poisson_log_probability(count: wp.float64, rate: wp.float64) -> wp.float64:200"""Stable log mass for a non-negative integer and a positive Poisson mean.201202Loader (2002), equations (4), (6), (7): express the mass through Stirling203error and deviance before evaluating it. Subtracting count*log(rate) and204lgamma(count+1) directly discards useful digits near a large mean.205https://www.r-project.org/doc/reports/CLoader-dbinom-2002.pdf206"""207if count == wp.float64(0.0):208return -rate209if count < wp.float64(16.0):210return -rate + count * wp.log(rate) - log_gamma(count + wp.float64(1.0))211inverse = wp.float64(1.0) / count212square = inverse * inverse213# Cast both operands: casting a quotient lets Warp divide in FP32 first.214# Six Bernoulli terms; at count >= 16 the next term bounds the absolute215# truncation error by 7/(1092*16**13) < 1.5e-18 before FP64 rounding.216correction = inverse * (217(wp.float64(1.0) / wp.float64(12.0))218- square219* (220(wp.float64(1.0) / wp.float64(360.0))221- square222* (223(wp.float64(1.0) / wp.float64(1260.0))224- square225* (226(wp.float64(1.0) / wp.float64(1680.0))227- square228* (229(wp.float64(1.0) / wp.float64(1188.0))230- square * (wp.float64(691.0) / wp.float64(360360.0))231)232)233)234)235)236difference = count - rate237deviance = wp.float64(0.0)238if wp.abs(difference) < wp.float64(0.1) * (count + rate):239ratio = difference / (count + rate)240ratio_squared = ratio * ratio241deviance = difference * ratio242term = wp.float64(2.0) * count * ratio243# |ratio| < .1 gives geometric convergence; 32 terms make the244# omitted relative tail smaller than binary64 precision throughout.245for order in range(1, 33):246term = term * ratio_squared247updated = deviance + term / wp.float64(2 * order + 1)248if updated == deviance:249break250deviance = updated251else:252deviance = count * wp.log(count / rate) - difference253return (254-wp.float64(0.5) * wp.log(wp.float64(6.283185307179586476925286766559) * count)255- correction256- deviance257)258259260# region book:detector-poisson-observation261@wp.func262def poisson_draw(263rate: wp.float64,264seed: wp.uint64,265identity: wp.uint64,266domain: wp.uint32,267budget: int,268status: wp.array(dtype=wp.int32),269) -> wp.uint64:270if rate == wp.float64(0.0):271return wp.uint64(0)272if not wp.isfinite(rate) or rate < wp.float64(0.0) or rate > wp.float64(1.0e9):273wp.atomic_or(status, 0, 1)274return wp.uint64(0)275if rate < wp.float64(10.0):276u = uniform_pair(seed, identity, wp.uint32(0), domain)[0]277probability = wp.exp(-rate)278cumulative = probability279count = int(0) # noqa: UP018, RUF046 - mutable Warp loop variable280# One inverse-CDF proposal. Its recurrence bound is independent of281# draw_budget: a small rate does not imply a bounded Poisson count.282# At rate < 10 the tail beyond 128 is far below the 2^-33 closest283# approach of our open-interval 32-bit uniform to one.284while u > cumulative and count < 128:285count = count + 1286probability = probability * rate / wp.float64(count)287cumulative = cumulative + probability288if u <= cumulative:289return wp.uint64(count)290else:291# Hörmann's transformed rejection with squeeze (PTRS), 1993,292# doi:10.1016/0167-6687(93)90997-4. All acceptance arithmetic is FP64.293root = wp.sqrt(rate)294b = wp.float64(0.931) + wp.float64(2.53) * root295a = wp.float64(-0.059) + wp.float64(0.02483) * b296inverse_alpha = wp.float64(1.1239) + wp.float64(1.1328) / (b - wp.float64(3.4))297squeeze = wp.float64(0.9277) - wp.float64(3.6224) / (b - wp.float64(2.0))298for trial in range(budget):299uniforms = uniform_pair(seed, identity, wp.uint32(trial), domain)300u = uniforms[0] - wp.float64(0.5)301v = uniforms[1]302distance = wp.float64(0.5) - wp.abs(u)303candidate = wp.floor((wp.float64(2.0) * a / distance + b) * u + rate + wp.float64(0.43))304if candidate >= wp.float64(0.0):305if distance >= wp.float64(0.07) and v <= squeeze:306return wp.uint64(candidate)307if not (distance < wp.float64(0.013) and v > distance):308lhs = wp.log(v * inverse_alpha / (a / (distance * distance) + b))309rhs = poisson_log_probability(candidate, rate)310if lhs <= rhs:311return wp.uint64(candidate)312# A bounded failure is an invalid realisation, never an observed zero.313wp.atomic_or(status, 0, 4)314return wp.uint64(0)315316317@wp.kernel318def poisson_counts(319mean: wp.array(dtype=wp.float32),320seed: wp.uint64,321observation: wp.uint32,322pixel_offset: wp.uint32,323energy_offset: wp.uint32,324budget: int,325output: wp.array(dtype=wp.uint64),326status: wp.array(dtype=wp.int32),327):328p = wp.tid()329identity = (wp.uint64(observation) << wp.uint64(32)) | wp.uint64(pixel_offset + wp.uint32(p))330output[p] = poisson_draw(331wp.float64(mean[p]), seed, identity, wp.uint32(1) + energy_offset, budget, status332)333334335@cache336def get_compound_poisson(energies: int, shared_scores: bool):337@wp.kernel(module="unique", module_options=STRICT)338def compound(339rates: wp.array(dtype=wp.float32),340scores: wp.array(dtype=wp.float32),341pixels: int,342seed: wp.uint64,343observation: wp.uint32,344pixel_offset: wp.uint32,345energy_offset: wp.uint32,346budget: int,347output: wp.array(dtype=wp.float32),348status: wp.array(dtype=wp.int32),349):350p = wp.tid()351identity = (wp.uint64(observation) << wp.uint64(32)) | wp.uint64(352pixel_offset + wp.uint32(p)353)354total = wp.float64(0.0)355compensation = wp.float64(0.0)356for energy in range(energies):357si = energy * pixels + p358if wp.static(shared_scores):359si = energy360count = poisson_draw(361wp.float64(rates[energy * pixels + p]),362seed,363identity,364wp.uint32(1) + energy_offset + wp.uint32(energy),365budget,366status,367)368corrected = wp.float64(count) * wp.float64(scores[si]) - compensation369updated = total + corrected370compensation = (updated - total) - corrected371total = updated372output[p] = checked_store(total, status)373374return compound375376377# endregion book:detector-poisson-observation378379380@wp.kernel381def gaussian_read_noise(382signal: wp.array(dtype=wp.float32),383sigma: wp.float64,384seed: wp.uint64,385observation: wp.uint32,386pixel_offset: wp.uint32,387output: wp.array(dtype=wp.float32),388status: wp.array(dtype=wp.int32),389):390p = wp.tid()391identity = (wp.uint64(observation) << wp.uint64(32)) | wp.uint64(pixel_offset + wp.uint32(p))392uniforms = uniform_pair(seed, identity, wp.uint32(0), wp.uint32(0x7FFFFFFF))393normal = wp.sqrt(-wp.float64(2.0) * wp.log(uniforms[0])) * wp.cos(394wp.float64(6.283185307179586476925286766559) * uniforms[1]395)396output[p] = checked_store(wp.float64(signal[p]) + sigma * normal, status)397