integrate_module 0.99.1__tar.gz → 0.99.3__tar.gz
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.
- {integrate_module-0.99.1/integrate_module.egg-info → integrate_module-0.99.3}/PKG-INFO +2 -1
- {integrate_module-0.99.1 → integrate_module-0.99.3}/integrate/__init__.py +2 -1
- {integrate_module-0.99.1 → integrate_module-0.99.3}/integrate/integrate.py +145 -7
- {integrate_module-0.99.1 → integrate_module-0.99.3}/integrate/integrate_io.py +20 -1
- {integrate_module-0.99.1 → integrate_module-0.99.3}/integrate/integrate_plot.py +384 -440
- {integrate_module-0.99.1 → integrate_module-0.99.3}/integrate/integrate_query.py +17 -3
- {integrate_module-0.99.1 → integrate_module-0.99.3}/integrate/integrate_timing_cli.py +5 -3
- {integrate_module-0.99.1 → integrate_module-0.99.3/integrate_module.egg-info}/PKG-INFO +2 -1
- {integrate_module-0.99.1 → integrate_module-0.99.3}/integrate_module.egg-info/requires.txt +1 -0
- {integrate_module-0.99.1 → integrate_module-0.99.3}/pyproject.toml +2 -1
- {integrate_module-0.99.1 → integrate_module-0.99.3}/LICENSE +0 -0
- {integrate_module-0.99.1 → integrate_module-0.99.3}/README.md +0 -0
- {integrate_module-0.99.1 → integrate_module-0.99.3}/integrate/gex.py +0 -0
- {integrate_module-0.99.1 → integrate_module-0.99.3}/integrate/integrate_borehole.py +0 -0
- {integrate_module-0.99.1 → integrate_module-0.99.3}/integrate/integrate_hdf5_info_cli.py +0 -0
- {integrate_module-0.99.1 → integrate_module-0.99.3}/integrate/integrate_rejection.py +0 -0
- {integrate_module-0.99.1 → integrate_module-0.99.3}/integrate/integrate_rejection_cli.py +0 -0
- {integrate_module-0.99.1 → integrate_module-0.99.3}/integrate/integrate_rejection_jax.py +0 -0
- {integrate_module-0.99.1 → integrate_module-0.99.3}/integrate/integrate_www_cli.py +0 -0
- {integrate_module-0.99.1 → integrate_module-0.99.3}/integrate_module.egg-info/SOURCES.txt +0 -0
- {integrate_module-0.99.1 → integrate_module-0.99.3}/integrate_module.egg-info/dependency_links.txt +0 -0
- {integrate_module-0.99.1 → integrate_module-0.99.3}/integrate_module.egg-info/entry_points.txt +0 -0
- {integrate_module-0.99.1 → integrate_module-0.99.3}/integrate_module.egg-info/top_level.txt +0 -0
- {integrate_module-0.99.1 → integrate_module-0.99.3}/setup.cfg +0 -0
- {integrate_module-0.99.1 → integrate_module-0.99.3}/tests/test_likelihood_multinomial.py +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: integrate_module
|
|
3
|
-
Version: 0.99.
|
|
3
|
+
Version: 0.99.3
|
|
4
4
|
Summary: Localized probabilistic data integration
|
|
5
5
|
Author-email: Thomas Mejer Hansen <tmeha@geo.au.dk>
|
|
6
6
|
License: MIT
|
|
@@ -30,6 +30,7 @@ Requires-Dist: pyvista
|
|
|
30
30
|
Requires-Dist: litellm
|
|
31
31
|
Requires-Dist: streamlit
|
|
32
32
|
Requires-Dist: libaarhusxyz
|
|
33
|
+
Requires-Dist: jax
|
|
33
34
|
Provides-Extra: dev
|
|
34
35
|
Requires-Dist: pytest; extra == "dev"
|
|
35
36
|
Requires-Dist: black; extra == "dev"
|
|
@@ -18,8 +18,9 @@ from integrate.integrate_rejection import create_shared_memory
|
|
|
18
18
|
from integrate.integrate_rejection import compute_hypothesis_probability
|
|
19
19
|
|
|
20
20
|
# Import other functions from main module
|
|
21
|
-
from integrate.integrate import integrate_update_prior_attributes
|
|
21
|
+
from integrate.integrate import integrate_update_prior_attributes
|
|
22
22
|
from integrate.integrate import integrate_posterior_stats
|
|
23
|
+
from integrate.integrate import prior_set
|
|
23
24
|
from integrate.integrate import logl_T_est
|
|
24
25
|
from integrate.integrate import lu_post_sample_logl
|
|
25
26
|
from integrate.integrate import prior_data
|
|
@@ -287,6 +287,131 @@ def integrate_update_prior_attributes(f_prior_h5, **kwargs):
|
|
|
287
287
|
|
|
288
288
|
|
|
289
289
|
|
|
290
|
+
def prior_set(f_prior_h5, im=1, **kwargs):
|
|
291
|
+
"""
|
|
292
|
+
Update attributes of a model parameter in a prior HDF5 file.
|
|
293
|
+
|
|
294
|
+
Parameters
|
|
295
|
+
----------
|
|
296
|
+
f_prior_h5 : str
|
|
297
|
+
Path to the prior HDF5 file to update.
|
|
298
|
+
im : int
|
|
299
|
+
Model parameter index (e.g., 1 for /M1, 2 for /M2, default is 1).
|
|
300
|
+
name : str, optional
|
|
301
|
+
Display name of the model parameter (e.g., 'Resistivity', 'Lithology').
|
|
302
|
+
class_name : list of str, optional
|
|
303
|
+
Class names for discrete parameters. If fewer names than existing classes
|
|
304
|
+
are given, only the first N entries are updated and the rest are kept.
|
|
305
|
+
color : list of str or RGBA tuples, optional
|
|
306
|
+
Per-class colors (e.g., ['blue', 'red']). Converted to RGBA and stored
|
|
307
|
+
as the ``cmap`` attribute. If fewer colors than existing classes are given,
|
|
308
|
+
only the first N entries are updated and the rest are kept.
|
|
309
|
+
class_id : array-like, optional
|
|
310
|
+
Class ID values for discrete parameters.
|
|
311
|
+
clim : list of float, optional
|
|
312
|
+
Color scale limits [min, max].
|
|
313
|
+
is_discrete : int or bool, optional
|
|
314
|
+
Whether the parameter is discrete (1) or continuous (0).
|
|
315
|
+
cmap : str or array-like, optional
|
|
316
|
+
Colormap. Either a matplotlib colormap name (colors are sampled to match
|
|
317
|
+
the number of classes) or an RGBA array of shape (4, N).
|
|
318
|
+
|
|
319
|
+
Examples
|
|
320
|
+
--------
|
|
321
|
+
>>> ig.prior_set('PRIOR.h5', im=1, name='Resistivity', clim=[1, 2600])
|
|
322
|
+
>>> ig.prior_set('PRIOR.h5', im=2, class_name=['Sand', 'Clay', 'Gravel'])
|
|
323
|
+
>>> ig.prior_set('PRIOR.h5', im=2, class_name=['inside valley', 'outside valley'],
|
|
324
|
+
... color=['blue', 'red'])
|
|
325
|
+
"""
|
|
326
|
+
import matplotlib.colors as mcolors
|
|
327
|
+
import matplotlib
|
|
328
|
+
|
|
329
|
+
def _set_attr(ds, key, value):
|
|
330
|
+
# Delete first to avoid h5py silently failing on dtype/shape changes
|
|
331
|
+
if key in ds.attrs:
|
|
332
|
+
del ds.attrs[key]
|
|
333
|
+
ds.attrs[key] = value
|
|
334
|
+
|
|
335
|
+
def _decode(v):
|
|
336
|
+
return v.decode('utf-8') if isinstance(v, bytes) else str(v)
|
|
337
|
+
|
|
338
|
+
Mstr = 'M%d' % im
|
|
339
|
+
|
|
340
|
+
if not os.path.exists(f_prior_h5):
|
|
341
|
+
print('prior_set: File %s does not exist (will not create)' % f_prior_h5)
|
|
342
|
+
return
|
|
343
|
+
|
|
344
|
+
with h5py.File(f_prior_h5, 'a') as f:
|
|
345
|
+
if Mstr not in f:
|
|
346
|
+
print('prior_set: %s not found in %s' % (Mstr, f_prior_h5))
|
|
347
|
+
return
|
|
348
|
+
|
|
349
|
+
ds = f[Mstr]
|
|
350
|
+
|
|
351
|
+
if 'name' in kwargs:
|
|
352
|
+
_set_attr(ds, 'name', str(kwargs['name']))
|
|
353
|
+
print('prior_set: %s/name = %r' % (Mstr, kwargs['name']))
|
|
354
|
+
|
|
355
|
+
if 'class_name' in kwargs:
|
|
356
|
+
new_names = [str(n) for n in kwargs['class_name']]
|
|
357
|
+
if 'class_name' in ds.attrs:
|
|
358
|
+
existing = [_decode(n) for n in ds.attrs['class_name']]
|
|
359
|
+
for i, n in enumerate(new_names):
|
|
360
|
+
if i < len(existing):
|
|
361
|
+
existing[i] = n
|
|
362
|
+
else:
|
|
363
|
+
existing.append(n)
|
|
364
|
+
final = existing
|
|
365
|
+
else:
|
|
366
|
+
final = new_names
|
|
367
|
+
_set_attr(ds, 'class_name', [str(n) for n in final])
|
|
368
|
+
print('prior_set: %s/class_name = %s' % (Mstr, list(ds.attrs['class_name'])))
|
|
369
|
+
|
|
370
|
+
if 'color' in kwargs:
|
|
371
|
+
new_colors = [np.array(mcolors.to_rgba(c)) for c in kwargs['color']]
|
|
372
|
+
n_new = len(new_colors)
|
|
373
|
+
if 'cmap' in ds.attrs:
|
|
374
|
+
existing_rgba = ds.attrs['cmap'].T # (N, 4)
|
|
375
|
+
existing_list = [existing_rgba[i] for i in range(len(existing_rgba))]
|
|
376
|
+
for i, c in enumerate(new_colors):
|
|
377
|
+
if i < len(existing_list):
|
|
378
|
+
existing_list[i] = c
|
|
379
|
+
else:
|
|
380
|
+
existing_list.append(c)
|
|
381
|
+
rgba = np.array(existing_list)
|
|
382
|
+
else:
|
|
383
|
+
rgba = np.array(new_colors)
|
|
384
|
+
_set_attr(ds, 'cmap', rgba.T) # stored as (4, N)
|
|
385
|
+
print('prior_set: %s/cmap updated from color list (%d colors)' % (Mstr, n_new))
|
|
386
|
+
|
|
387
|
+
if 'class_id' in kwargs:
|
|
388
|
+
_set_attr(ds, 'class_id', np.array(kwargs['class_id']))
|
|
389
|
+
print('prior_set: %s/class_id = %s' % (Mstr, ds.attrs['class_id']))
|
|
390
|
+
|
|
391
|
+
if 'clim' in kwargs:
|
|
392
|
+
_set_attr(ds, 'clim', np.array(kwargs['clim'], dtype=float))
|
|
393
|
+
print('prior_set: %s/clim = %s' % (Mstr, ds.attrs['clim']))
|
|
394
|
+
|
|
395
|
+
if 'is_discrete' in kwargs:
|
|
396
|
+
_set_attr(ds, 'is_discrete', int(kwargs['is_discrete']))
|
|
397
|
+
print('prior_set: %s/is_discrete = %d' % (Mstr, ds.attrs['is_discrete']))
|
|
398
|
+
|
|
399
|
+
if 'cmap' in kwargs:
|
|
400
|
+
cmap_val = kwargs['cmap']
|
|
401
|
+
if isinstance(cmap_val, str):
|
|
402
|
+
if 'class_id' in ds.attrs:
|
|
403
|
+
n_colors = len(ds.attrs['class_id'])
|
|
404
|
+
elif 'class_name' in ds.attrs:
|
|
405
|
+
n_colors = len(ds.attrs['class_name'])
|
|
406
|
+
else:
|
|
407
|
+
n_colors = 10
|
|
408
|
+
rgba = matplotlib.colormaps[cmap_val](np.linspace(0, 1, n_colors))
|
|
409
|
+
else:
|
|
410
|
+
rgba = np.array(cmap_val)
|
|
411
|
+
_set_attr(ds, 'cmap', rgba.T) # stored as (4, N)
|
|
412
|
+
print('prior_set: %s/cmap updated' % Mstr)
|
|
413
|
+
|
|
414
|
+
|
|
290
415
|
def integrate_posterior_stats(f_post_h5='POST.h5', ip_range=None, **kwargs):
|
|
291
416
|
"""
|
|
292
417
|
Compute posterior statistics for all model parameters in a POST HDF5 file.
|
|
@@ -3306,37 +3431,41 @@ def timing_compute(N_arr=[], Nproc_arr=[], backend='numpy', NcpuForward=0):
|
|
|
3306
3431
|
return file_out
|
|
3307
3432
|
|
|
3308
3433
|
|
|
3309
|
-
def timing_plot(f_timing=''):
|
|
3434
|
+
def timing_plot(f_timing='', fontsize=16):
|
|
3310
3435
|
"""
|
|
3311
3436
|
Generate comprehensive timing analysis plots from benchmark results.
|
|
3312
|
-
|
|
3437
|
+
|
|
3313
3438
|
This function creates multiple plots analyzing the performance characteristics
|
|
3314
3439
|
of the INTEGRATE workflow across different dataset sizes and processor counts.
|
|
3315
|
-
|
|
3440
|
+
|
|
3316
3441
|
Parameters
|
|
3317
3442
|
----------
|
|
3318
3443
|
f_timing : str
|
|
3319
3444
|
Path to NPZ file containing timing benchmark results from timing_compute().
|
|
3320
|
-
|
|
3445
|
+
fontsize : int, optional
|
|
3446
|
+
Base font size for all plot text (labels, titles, ticks, legends).
|
|
3447
|
+
Default is 16.
|
|
3448
|
+
|
|
3321
3449
|
Returns
|
|
3322
3450
|
-------
|
|
3323
3451
|
None
|
|
3324
3452
|
Saves multiple PNG files with timing analysis plots.
|
|
3325
|
-
|
|
3453
|
+
|
|
3326
3454
|
Notes
|
|
3327
3455
|
-----
|
|
3328
3456
|
Generated plots include:
|
|
3329
3457
|
- Total execution time vs processors and dataset size
|
|
3330
|
-
- Forward modeling performance and speedup analysis
|
|
3458
|
+
- Forward modeling performance and speedup analysis
|
|
3331
3459
|
- Rejection sampling performance and scaling
|
|
3332
3460
|
- Posterior statistics computation performance
|
|
3333
3461
|
- Cumulative time breakdowns for different processor counts
|
|
3334
3462
|
- Comparisons with traditional least squares and MCMC methods
|
|
3335
|
-
|
|
3463
|
+
|
|
3336
3464
|
The function handles missing data gracefully and includes reference lines
|
|
3337
3465
|
for linear scaling to assess parallel efficiency.
|
|
3338
3466
|
"""
|
|
3339
3467
|
import numpy as np
|
|
3468
|
+
import matplotlib as mpl
|
|
3340
3469
|
import matplotlib.pyplot as plt
|
|
3341
3470
|
|
|
3342
3471
|
def safe_show():
|
|
@@ -3351,6 +3480,15 @@ def timing_plot(f_timing=''):
|
|
|
3351
3480
|
else:
|
|
3352
3481
|
print('Plotting timing results from %s' % f_timing)
|
|
3353
3482
|
|
|
3483
|
+
mpl.rcParams.update({
|
|
3484
|
+
'font.size': fontsize,
|
|
3485
|
+
'axes.labelsize': fontsize,
|
|
3486
|
+
'axes.titlesize': fontsize + 2,
|
|
3487
|
+
'xtick.labelsize': fontsize - 2,
|
|
3488
|
+
'ytick.labelsize': fontsize - 2,
|
|
3489
|
+
'legend.fontsize': fontsize - 2,
|
|
3490
|
+
})
|
|
3491
|
+
|
|
3354
3492
|
# file_out is f_timing, without file extension
|
|
3355
3493
|
file_out = f_timing.split('.')[0]
|
|
3356
3494
|
|
|
@@ -2967,6 +2967,10 @@ def get_case_data(case='DAUGAARD', loadAll=False, loadType='', filelist=None, **
|
|
|
2967
2967
|
print("filelist to download:")
|
|
2968
2968
|
print(filelist)
|
|
2969
2969
|
|
|
2970
|
+
# Direct-download endpoint for the shared data. Use ERDA's 'share_redirect'
|
|
2971
|
+
# form (serves raw file bytes, returns proper 404s), NOT the 'sharelink'
|
|
2972
|
+
# browse link, which returns an HTML directory page with status 200 for any
|
|
2973
|
+
# path and would be saved as a corrupt file.
|
|
2970
2974
|
urlErda = 'https://anon.erda.au.dk/share_redirect/dxOLKDtoul'
|
|
2971
2975
|
urlErdaCase = '%s/%s' % (urlErda,case)
|
|
2972
2976
|
from tqdm import tqdm
|
|
@@ -2976,7 +2980,22 @@ def get_case_data(case='DAUGAARD', loadAll=False, loadType='', filelist=None, **
|
|
|
2976
2980
|
if showInfo>-1:
|
|
2977
2981
|
print('--> Got data for case: %s' % case)
|
|
2978
2982
|
|
|
2979
|
-
|
|
2983
|
+
# Return the local filename for each requested file, but only if it was
|
|
2984
|
+
# actually obtained (downloaded or already present locally). download_file
|
|
2985
|
+
# saves to download_dir='.' using the basename, so we check for that here.
|
|
2986
|
+
# Files that could not be obtained (e.g. missing on the remote server)
|
|
2987
|
+
# return an empty string, so callers can test `if len(name) == 0`.
|
|
2988
|
+
result = []
|
|
2989
|
+
for f in filelist:
|
|
2990
|
+
basename = f.replace('\\', '/').split('/')[-1]
|
|
2991
|
+
if os.path.exists(basename):
|
|
2992
|
+
result.append(basename)
|
|
2993
|
+
else:
|
|
2994
|
+
if showInfo>-1:
|
|
2995
|
+
print('File %s was not obtained (missing locally and on remote); returning empty string.' % basename)
|
|
2996
|
+
result.append('')
|
|
2997
|
+
|
|
2998
|
+
return result
|
|
2980
2999
|
|
|
2981
3000
|
|
|
2982
3001
|
|