PyOPIA 1.1.12__tar.gz → 2.0.0__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.
@@ -1,20 +1,21 @@
1
1
  Metadata-Version: 2.1
2
2
  Name: PyOPIA
3
- Version: 1.1.12
3
+ Version: 2.0.0
4
4
  Summary: A Python Ocean Particle Image Analysis toolbox.
5
5
  Home-page: https://github.com/sintef/pyopia
6
6
  Keywords: Ocean,Particles,Imaging,Measurement,Size distribution
7
7
  Author: Emlyn Davies
8
8
  Author-email: emlyn.davies@sintef.no
9
- Requires-Python: >=3.10,<3.11
9
+ Requires-Python: >=3.12,<4.0
10
10
  Classifier: Programming Language :: Python :: 3
11
- Classifier: Programming Language :: Python :: 3.10
11
+ Classifier: Programming Language :: Python :: 3.12
12
12
  Provides-Extra: classification
13
13
  Provides-Extra: classification-arm64
14
14
  Requires-Dist: cmocean (>=3.0.3,<4.0.0)
15
15
  Requires-Dist: dask (>=2024.8.1)
16
16
  Requires-Dist: flake8 (>=6.1.0,<7.0.0)
17
17
  Requires-Dist: gdown (>=4.7.1,<5.0.0)
18
+ Requires-Dist: h5netcdf (>=1.3.0)
18
19
  Requires-Dist: h5py (>=3.9.0,<4.0.0)
19
20
  Requires-Dist: imageio (>=2.31.3,<3.0.0)
20
21
  Requires-Dist: ipykernel (>=6.19.4)
@@ -22,6 +23,7 @@ Requires-Dist: jupyter-book (>=0.15.1,<0.16.0)
22
23
  Requires-Dist: matplotlib (>=3.7)
23
24
  Requires-Dist: myst-nb (>=0.17.2,<0.18.0)
24
25
  Requires-Dist: nbclient (==0.7)
26
+ Requires-Dist: nbconvert (>=7.16.4,<8.0.0)
25
27
  Requires-Dist: numpy (>=1.24.0,<2.0.0)
26
28
  Requires-Dist: pandas[computation] (>=2.1.1,<3.0.0)
27
29
  Requires-Dist: poetry-version-plugin (>=0.2.0,<0.3.0)
@@ -35,9 +37,8 @@ Requires-Dist: sphinx-copybutton (>=0.5.2,<0.6.0)
35
37
  Requires-Dist: sphinx-rtd-theme (>=0.5.0)
36
38
  Requires-Dist: sphinx-togglebutton (>=0.3.2,<0.4.0)
37
39
  Requires-Dist: sphinxcontrib-napoleon (>=0.7)
38
- Requires-Dist: tensorflow-cpu (==2.11.0) ; extra == "classification"
39
- Requires-Dist: tensorflow-io-gcs-filesystem (>=0.31.0) ; extra == "classification-arm64" or extra == "classification"
40
- Requires-Dist: tensorflow-macos (==2.11.0) ; (sys_platform == "darwin" and platform_machine == "arm64") and (extra == "classification-arm64")
40
+ Requires-Dist: tensorflow-cpu (>=2.16.2,<3.0.0) ; extra == "classification"
41
+ Requires-Dist: tensorflow-macos (>=2.16.2,<3.0.0) ; (sys_platform == "darwin" and platform_machine == "arm64") and (extra == "classification-arm64")
41
42
  Requires-Dist: toml (>=0.10.2,<0.11.0)
42
43
  Requires-Dist: tqdm (>=4.66.1,<5.0.0)
43
44
  Requires-Dist: typer[all] (>=0.9.0,<0.10.0)
@@ -143,12 +144,23 @@ or for arm/silicon systems:
143
144
  poetry install --extras "classification-arm64"
144
145
  ```
145
146
 
147
+ Note: If poetry spends ages resolving dependencies, you can install a development environment with pip, like this: `pip install -e ".[classification]"` or for arm/silicon: `pip install -e ".[classification-arm64]"`
148
+
149
+
146
150
  3. (optional) Run local tests:
147
151
 
148
152
  ```bash
149
153
  poetry run pytest
150
154
  ```
151
155
 
156
+ #### Version numbering
157
+
158
+ The version number of PyOPIA is split into three sections: MAJOR.MINOR.PATCH
159
+
160
+ * MAJOR: Changes in high-level pipeline use and/or data output that are not backwards-compatible.
161
+ * MINOR: New features that are backwards-compatible.
162
+ * PATCH: Backwards-compatible bug fixes or enhancements to existing functionality
163
+
152
164
  ## Build docs locally
153
165
 
154
166
  ```
@@ -94,12 +94,23 @@ or for arm/silicon systems:
94
94
  poetry install --extras "classification-arm64"
95
95
  ```
96
96
 
97
+ Note: If poetry spends ages resolving dependencies, you can install a development environment with pip, like this: `pip install -e ".[classification]"` or for arm/silicon: `pip install -e ".[classification-arm64]"`
98
+
99
+
97
100
  3. (optional) Run local tests:
98
101
 
99
102
  ```bash
100
103
  poetry run pytest
101
104
  ```
102
105
 
106
+ #### Version numbering
107
+
108
+ The version number of PyOPIA is split into three sections: MAJOR.MINOR.PATCH
109
+
110
+ * MAJOR: Changes in high-level pipeline use and/or data output that are not backwards-compatible.
111
+ * MINOR: New features that are backwards-compatible.
112
+ * PATCH: Backwards-compatible bug fixes or enhancements to existing functionality
113
+
103
114
  ## Build docs locally
104
115
 
105
116
  ```
@@ -0,0 +1 @@
1
+ __version__ = '2.0.0'
@@ -1,10 +1,7 @@
1
1
  '''
2
2
  Background correction module (inherited from PySilCam)
3
3
  '''
4
-
5
4
  import numpy as np
6
- from glob import glob
7
- from pyopia.pipeline import get_load_function
8
5
 
9
6
 
10
7
  def ini_background(bgfiles, load_function):
@@ -86,19 +83,21 @@ def correct_im_accurate(imbg, imraw):
86
83
  highlights if the background or raw images are very poorly obtained
87
84
 
88
85
  Args:
89
- imbg (uint8 or float64) : background averaged image
90
- imraw (uint8 or float64) : raw image
86
+ imbg (float64) : background averaged image
87
+ imraw (float64) : raw image
88
+ imbg (float64) : background averaged image
89
+ imraw (float64) : raw image
91
90
 
92
91
  Returns:
93
- imc (uint8 or float64) : corrected image, same type as input
92
+ im_corrected (float64) : corrected image, same type as input
94
93
  '''
95
94
 
96
- imc = np.float64(imraw) - np.float64(imbg)
97
- imc += (255 / 2 - np.percentile(imc, 50))
95
+ im_corrected = imraw - imbg
96
+ im_corrected += (1 / 2 - np.percentile(im_corrected, 50))
98
97
 
99
- imc += 255 - imc.max()
98
+ im_corrected += 1 - im_corrected.max()
100
99
 
101
- return imc
100
+ return im_corrected
102
101
 
103
102
 
104
103
  def correct_im_fast(imbg, imraw):
@@ -110,20 +109,18 @@ def correct_im_fast(imbg, imraw):
110
109
  highlights, especially if the background or raw images are not properly obtained
111
110
 
112
111
  Args:
113
- imbg (uint8) : background averaged image
114
- imraw (uint8) : raw image
112
+ imraw (float64) : raw image
113
+ imbg (float64) : background averaged image
115
114
 
116
115
  Returns:
117
- imc (uint8) : corrected image
116
+ im_corrected (float64) : corrected image
118
117
  '''
119
- imc = imraw - imbg
118
+ im_corrected = imraw - imbg
120
119
 
121
- imc += 215
122
- imc[imc < 0] = 0
123
- imc[imc > 255] = 255
124
- imc = np.uint8(imc)
120
+ im_corrected += 215/255
121
+ im_corrected = np.clip(im_corrected, 0, 1)
125
122
 
126
- return imc
123
+ return im_corrected
127
124
 
128
125
 
129
126
  def shift_and_correct(bgstack, imbg, imraw, stacklength, real_time_stats=False):
@@ -135,162 +132,26 @@ def shift_and_correct(bgstack, imbg, imraw, stacklength, real_time_stats=False):
135
132
 
136
133
  Args:
137
134
  bgstack (list) : list of all images in the background stack
138
- imbg (uint8) : background image
139
- imraw (uint8) : raw image
135
+ imbg (float64) : background image
136
+ imraw (float64) : raw image
140
137
  stacklength (int) : unsed int here - just there to maintain the same behaviour as
141
138
  shift_bgstack_fast()
142
139
  real_time_stats=False (Bool) : if True use fast functions, if False use accurate functions
143
140
 
144
141
  Returns:
145
142
  bgstack (list) : list of all images in the background stack
146
- imbg (uint8) : background averaged image
147
- imc (uint8) : corrected image
143
+ imbg (float64) : background averaged image
144
+ im_corrected (float64) : corrected image
148
145
  '''
149
146
 
150
147
  if real_time_stats:
151
- imc = correct_im_fast(imbg, imraw)
148
+ im_corrected = correct_im_fast(imbg, imraw)
152
149
  bgstack, imbg = shift_bgstack_fast(bgstack, imbg, imraw, stacklength)
153
150
  else:
154
- imc = correct_im_accurate(imbg, imraw)
151
+ im_corrected = correct_im_accurate(imbg, imraw)
155
152
  bgstack, imbg = shift_bgstack_accurate(bgstack, imbg, imraw, stacklength)
156
153
 
157
- return bgstack, imbg, imc
158
-
159
-
160
- def backgrounder(av_window, acquire, bad_lighting_limit=None,
161
- real_time_stats=False):
162
- '''
163
- Generator which interacts with acquire to return a corrected image
164
- given av_window number of frame to use in creating a moving background
165
-
166
- Args:
167
- av_window (int) : number of images to use in creating the background
168
- acquire (generator object) : acquire generator object created by the Acquire class
169
- bad_lighting_limit=None (int) : if a number is supplied it is used for throwing away raw images that have a
170
- standard deviation in colour which exceeds the given value
171
-
172
- Yields:
173
- timestamp (timestamp) : timestamp of when raw image was acquired
174
- imc (uint8) : corrected image ready for analysis or plotting
175
- imraw (uint8) : raw image
176
-
177
- Example:
178
-
179
- .. code-block:: python
180
-
181
- avwind = 10 # number of images used for background
182
- imgen = backgrounder(avwind,acquire,bad_lighting_limit) # setup generator
183
-
184
- n = 10 # acquire 10 images and correct them with a sliding background:
185
- for i in range(n):
186
- imc = next(imgen)
187
- print(i)
188
- '''
189
-
190
- # Set up initial background image stack
191
- bgstack, imbg = ini_background(av_window, acquire)
192
- stacklength = len(bgstack)
193
-
194
- # Aquire images, apply background correction and yield result
195
- for timestamp, imraw in acquire:
196
-
197
- if bad_lighting_limit is not None:
198
- bgstack_new, imbg_new, imc = shift_and_correct(bgstack, imbg,
199
- imraw, stacklength, real_time_stats)
200
-
201
- # basic check of image quality
202
- r = imc[:, :, 0]
203
- g = imc[:, :, 1]
204
- b = imc[:, :, 2]
205
- s = np.std([r, g, b])
206
- # ignore bad images
207
- if s <= bad_lighting_limit:
208
- bgstack = bgstack_new
209
- imbg = imbg_new
210
- yield timestamp, imc, imraw
211
- else:
212
- print('bad lighting, std={0}'.format(s))
213
- else:
214
- bgstack, imbg, imc = shift_and_correct(bgstack, imbg, imraw,
215
- stacklength, real_time_stats)
216
- yield timestamp, imc, imraw
217
-
218
-
219
- def subtract_background(imbg, imraw):
220
- ''' simple background substraction
221
-
222
- Returns
223
- -------
224
- np.array : image
225
- image corrected by simple subtraction (imraw - imbw)
226
- '''
227
- return imraw - imbg
228
-
229
-
230
- class CreateBackground():
231
- '''
232
- :class:`pyopia.pipeline` compatible class that calls: :func:`pyopia.background.ini_background`.
233
- This runs by default in the pipeline initial steps if named as 'createbackground'.
234
-
235
- Pipeline input data:
236
- --------------------
237
- :class:`pyopia.pipeline.Data`
238
-
239
- containing the following keys:
240
-
241
- :attr:`pyopia.pipeline.Data.raw_files`
242
-
243
- :attr:`pyopia.pipeline.Data.imc`
244
-
245
- :attr:`pyopia.pipeline.Data.bgstack`
246
-
247
- :attr:`pyopia.pipeline.Data.imraw`
248
-
249
- :attr:`pyopia.pipeline.Data.imbg`
250
-
251
- Parameters:
252
- -----------
253
- average_window : int
254
- number of images to use in the background image stack
255
-
256
- instrument_module: (str, optional)
257
- Defaults to 'imread' if not defined
258
- Other alternatives are: 'holo' or 'silcam' if you want to use the `load_image`functions
259
- implemented within the {mod}`pyopia.instrument` submodule.
260
-
261
- Returns:
262
- --------
263
- :class:`pyopia.pipeline.Data`
264
- containing the following new keys:
265
-
266
- :attr:`pyopia.pipeline.Data.bgstack`
267
-
268
- :attr:`pyopia.pipeline.Data.imbg`
269
-
270
- Example pipeline uses:
271
- ----------------------
272
-
273
- .. code-block:: python
274
-
275
- [steps.createbackground]
276
- pipeline_class = 'pyopia.background.CreateBackground'
277
- average_window = 10
278
- instrument_module = 'holo'
279
- '''
280
-
281
- def __init__(self, average_window, instrument_module='imread'):
282
- self.average_window = average_window
283
- self.load_function = get_load_function(instrument_module)
284
- pass
285
-
286
- def __call__(self, data):
287
- files = glob(data['raw_files'])
288
- bgfiles = files[:self.average_window]
289
- bgstack, imbg = ini_background(bgfiles, self.load_function)
290
-
291
- data['bgstack'] = bgstack
292
- data['imbg'] = imbg
293
- return data
154
+ return bgstack, imbg, im_corrected
294
155
 
295
156
 
296
157
  class CorrectBackgroundAccurate():
@@ -298,18 +159,18 @@ class CorrectBackgroundAccurate():
298
159
  :class:`pyopia.pipeline` compatible class that calls: :func:`pyopia.background.correct_im_accurate`
299
160
  and will shift the background using a moving average function if given.
300
161
 
162
+ The background stack and background image are created during the first 'average_window' (int) calls
163
+ to this class, and the skip_next_steps flag is set in the pipeline Data. No background correction
164
+ is performed during these steps.
165
+
301
166
  Pipeline input data:
302
167
  --------------------
303
168
  :class:`pyopia.pipeline.Data`
304
169
 
305
170
  containing the following keys:
306
171
 
307
- :attr:`pyopia.pipeline.Data.bgstack`
308
-
309
172
  :attr:`pyopia.pipeline.Data.imraw`
310
173
 
311
- :attr:`pyopia.pipeline.Data.imbg`
312
-
313
174
  Parameters:
314
175
  -----------
315
176
  bgshift_function : (string, optional)
@@ -320,18 +181,25 @@ class CorrectBackgroundAccurate():
320
181
 
321
182
  :func:`pyopia.background.shift_bgstack_fast`
322
183
 
184
+ average_window : int
185
+ number of images to use in the background image stack
186
+
187
+ image_source: (str, optional)
188
+ The key in Pipeline.data of the image to be background corrected.
189
+ Defaults to 'imraw'
190
+
323
191
  Returns:
324
192
  --------
325
193
  :class:`pyopia.pipeline.Data`
326
194
  containing the following new keys:
327
195
 
328
- :attr:`pyopia.pipeline.Data.imc`
196
+ :attr:`pyopia.pipeline.Data.im_corrected`
197
+ :attr:`pyopia.pipeline.Data.im_corrected`
329
198
 
330
199
  :attr:`pyopia.pipeline.Data.bgstack`
331
200
 
332
201
  :attr:`pyopia.pipeline.Data.imbg`
333
202
 
334
-
335
203
  Example pipeline uses:
336
204
  ----------------------
337
205
  Apply moving average using :func:`pyopia.background.shift_bgstack_accurate` :
@@ -341,6 +209,7 @@ class CorrectBackgroundAccurate():
341
209
  [steps.correctbackground]
342
210
  pipeline_class = 'pyopia.background.CorrectBackgroundAccurate'
343
211
  bgshift_function = 'accurate'
212
+ average_window = 5
344
213
 
345
214
  Apply static background correction:
346
215
 
@@ -349,18 +218,42 @@ class CorrectBackgroundAccurate():
349
218
  [steps.correctbackground]
350
219
  pipeline_class = 'pyopia.background.CorrectBackgroundAccurate'
351
220
  bgshift_function = 'pass'
352
-
221
+ average_window = 5
353
222
 
354
223
  If you do not want to do background correction, leave this step out of the pipeline.
355
224
  Then you could use :class:`pyopia.pipeline.CorrectBackgroundNone` if you need to instead.
356
225
  '''
357
226
 
358
- def __init__(self, bgshift_function='pass'):
227
+ def __init__(self, bgshift_function='pass', average_window=1, image_source='imraw'):
359
228
  self.bgshift_function = bgshift_function
360
- pass
229
+ self.average_window = average_window
230
+ self.image_source = image_source
231
+
232
+ def _build_background_step(self, data):
233
+ '''Add one layer to the background stack from the raw image in data pipeline, and update the background image.'''
234
+ if 'bgstack' not in data:
235
+ data['bgstack'] = []
236
+
237
+ init_complete = True
238
+ if len(data['bgstack']) < self.average_window:
239
+ data['bgstack'].append(data[self.image_source])
240
+ data['imbg'] = np.mean(data['bgstack'], axis=0)
241
+ init_complete = False
242
+
243
+ return init_complete
361
244
 
362
245
  def __call__(self, data):
363
- data['imc'] = correct_im_accurate(data['imbg'], data['imraw'])
246
+ # Initialize the background while required bgstack size not reached
247
+ init_complete = self._build_background_step(data)
248
+
249
+ # If we are still building the bgstack, return without doing image correction and bgstack update
250
+ if not init_complete:
251
+ # Flag to the pipeline that remaining steps should be skipped since we are still building the background
252
+ data['skip_next_steps'] = True
253
+ return data
254
+
255
+ data['im_corrected'] = correct_im_accurate(data['imbg'], data[self.image_source])
256
+ data['im_corrected'] = correct_im_accurate(data['imbg'], data[self.image_source])
364
257
 
365
258
  match self.bgshift_function:
366
259
  case 'pass':
@@ -368,18 +261,19 @@ class CorrectBackgroundAccurate():
368
261
  case 'accurate':
369
262
  data['bgstack'], data['imbg'] = shift_bgstack_accurate(data['bgstack'],
370
263
  data['imbg'],
371
- data['imraw'])
264
+ data[self.image_source])
372
265
  case 'fast':
373
266
  data['bgstack'], data['imbg'] = shift_bgstack_fast(data['bgstack'],
374
267
  data['imbg'],
375
- data['imraw'])
268
+ data[self.image_source])
376
269
  return data
377
270
 
378
271
 
379
272
  class CorrectBackgroundNone():
380
273
  '''
381
274
  :class:`pyopia.pipeline` compatible class for use when no background correction is required.
382
- This simply makes `data['imc'] = data['imraw'] in the pipeline.
275
+ This simply makes `data['im_corrected'] = data['imraw'] in the pipeline.
276
+ This simply makes `data['im_corrected'] = data['imraw'] in the pipeline.
383
277
 
384
278
  Pipeline input data:
385
279
  --------------------
@@ -398,7 +292,7 @@ class CorrectBackgroundNone():
398
292
  :class:`pyopia.pipeline.Data`
399
293
  containing the following new keys:
400
294
 
401
- :attr:`pyopia.pipeline.Data.imc`
295
+ :attr:`pyopia.pipeline.Data.im_corrected`
402
296
 
403
297
 
404
298
  Example pipeline uses:
@@ -416,6 +310,6 @@ class CorrectBackgroundNone():
416
310
  pass
417
311
 
418
312
  def __call__(self, data):
419
- data['imc'] = data['imraw']
313
+ data['im_corrected'] = data['imraw']
420
314
 
421
315
  return data
@@ -0,0 +1,182 @@
1
+ '''
2
+ Module containing tools for classifying particle ROIs
3
+ '''
4
+
5
+ import os
6
+ import numpy as np
7
+ import pandas as pd
8
+
9
+ import logging
10
+ logger = logging.getLogger()
11
+
12
+ # import tensorflow here. It must be imported on the processor where it will be used!
13
+ # import is therefore here instead of at the top of file.
14
+ # consider # noqa: E(?) for flake8 / linting
15
+ try:
16
+ from tensorflow import keras
17
+ import tensorflow as tf
18
+ except ImportError:
19
+ info_str = 'ERROR: Could not import Keras. Classify will not work'
20
+ info_str += ' until you install tensorflow.\n'
21
+ info_str += 'Use: pip install pyopia[classification]\n'
22
+ info_str += ' or: pip install pyopia[classification-arm64]'
23
+ info_str += ' for tensorflow-macos (silicon chips)'
24
+ raise ImportError(info_str)
25
+
26
+
27
+ class Classify():
28
+ '''
29
+ A classifier class for PyOPIA workflow.
30
+ This is intended as a parent class that can be used as a template for flexible classification methods
31
+
32
+ Args:
33
+ model_path=model_path (str) : path to particle-classifier e.g.
34
+ '/testdata/model_name/particle_classifier.h5'
35
+
36
+ Example:
37
+
38
+ .. code-block:: python
39
+
40
+ cl = Classify(model_path='/testdata/model_name/particle_classifier.h5')
41
+
42
+ prediction = cl.proc_predict(roi) # roi is an image roi to be classified
43
+
44
+ Note that :meth:`Classify.load_model()`
45
+ is run when the :class:`Classify` class is initialised.
46
+ If this is used in combination with multiprocessing then the model must be loaded
47
+ on the process where it will be used and not passed between processers
48
+ (i.e. cl must be initialised on that process).
49
+
50
+ The config setup looks like this:
51
+
52
+ .. code-block:: python
53
+
54
+ [steps.classifier]
55
+ pipeline_class = 'pyopia.classify.Classify'
56
+ model_path = 'keras_model.h5' # path to trained nn model
57
+
58
+ If `[steps.classifier]`is not defined, the classification will be skipped and no probabilities reported.
59
+
60
+ If you want to use an example trained model for SilCam data
61
+ (no guarantee of accuracy for other applications), you can get it using `exampledata`
62
+ within the notebooks folder (https://github.com/SINTEF/pyopia/blob/main/notebooks/exampledata.py):
63
+
64
+ .. code-block:: python
65
+
66
+ model_path = exampledata.get_example_model()
67
+
68
+ '''
69
+ def __init__(self, model_path=None):
70
+ self.model_path = model_path
71
+ self.load_model()
72
+
73
+ # Enable this to perform whitebalance correction in the preprocessing step
74
+ self.correct_whitebalance = False
75
+
76
+ def __call__(self):
77
+ return self
78
+
79
+ def load_model(self):
80
+ '''
81
+ Load a trained Keras model into the Classify class.
82
+
83
+ self.model (tf model object) : loaded Keras model
84
+ self.class_names (list) : names for the model output classes
85
+ '''
86
+ model_path = self.model_path
87
+
88
+ os.environ['TF_CPP_MIN_LOG_LEVEL'] = '2'
89
+ keras.backend.clear_session()
90
+
91
+ # Instantiate Keras model from file
92
+ path, filename = os.path.split(model_path)
93
+ self.model = keras.models.load_model(model_path)
94
+
95
+ # Try to create model output class name list from last model layer name
96
+ class_labels = None
97
+ try:
98
+ class_labels = self.model.layers[-1].name.split('.')
99
+ except: # noqa E722
100
+ logger.info('Could not get class names from model layer name, reverting to old method with header file.')
101
+
102
+ # If we could not create correct class names above, revert to old header file method
103
+ expected_class_number = self.model.layers[-1].output.shape[1]
104
+ if class_labels is None or len(class_labels) != expected_class_number:
105
+ header = pd.read_csv(os.path.join(path, 'header.tfl.txt'))
106
+ class_labels = header.columns
107
+
108
+ self.class_labels = class_labels
109
+ logger.info(self.class_labels)
110
+
111
+ def preprocessing(self, img_input):
112
+ '''
113
+ Preprocess ROI ready for prediction. example here based on the pysilcam network setup
114
+
115
+ Args:
116
+ img_input (float) : a particle ROI before preprocessing with range 0-1
117
+
118
+ Returns:
119
+ img_preprocessed (float) : a particle ROI with range 0.-255., corrected and preprocessed, ready for prediction
120
+ '''
121
+
122
+ whitebalanced = np.copy(img_input).astype(np.float64)
123
+
124
+ # Do white-balance correction as a per-channel histogram shift
125
+ if self.correct_whitebalance:
126
+ p = 99
127
+ for c in range(3):
128
+ whitebalanced[:, :, c] += (p/100) - np.percentile(whitebalanced[:, :, c], p)
129
+ whitebalanced[whitebalanced > 1] = 1
130
+ whitebalanced[whitebalanced < 0] = 0
131
+
132
+ # convert back to 0-255 scaling (because of this layer in the network:
133
+ # layers.Rescaling(1./255, input_shape=(img_height, img_width, 3)))
134
+ # This is useful because it allows training to use tf.keras.utils.image_dataset_from_directory,
135
+ # which loads images in 0-255 range
136
+ img = keras.utils.img_to_array(whitebalanced * 255)
137
+
138
+ # Get config for image resizing from the model
139
+ _, img_height, img_width, _ = self.model.get_config()['layers'][0]['config']['batch_shape']
140
+ pad_to_aspect_ratio = getattr(self.model.layers[0], 'pad_to_aspect_ratio', False)
141
+
142
+ # resize to match the dimentions expected by the network
143
+ img = tf.image.resize(img, [img_height, img_width],
144
+ method=tf.image.ResizeMethod.BILINEAR,
145
+ preserve_aspect_ratio=pad_to_aspect_ratio)
146
+
147
+ img_array = tf.keras.utils.img_to_array(img)
148
+ img_preprocessed = tf.expand_dims(img_array, 0) # Create a batch
149
+ return img_preprocessed
150
+
151
+ @tf.function
152
+ def predict(self, img_preprocessed):
153
+ '''
154
+ Use tensorflow model to classify particles. example here based on the pysilcam network setup.
155
+
156
+ Args:
157
+ img_preprocessed (float) : a particle ROI arry, corrected and preprocessed using :meth:`Classify.preprocessing`,
158
+ ready for prediction using :meth:`Classify.predict`
159
+
160
+ Returns:
161
+ prediction (array) : the probability of the roi belonging to each class
162
+ '''
163
+
164
+ prediction = self.model(img_preprocessed, training=False)
165
+ prediction = tf.nn.softmax(prediction[0])
166
+ return prediction
167
+
168
+ def proc_predict(self, img_input):
169
+ '''
170
+ Run pre-processing (:meth:`Classify.preprocessing`) and prediction (:meth:`Classify.predict`)
171
+ using tensorflow model to classify particles. example here based on the pysilcam network setup.
172
+
173
+ Args:
174
+ img_input (float) : a particle ROI with range 0-1 before preprocessing
175
+
176
+ Returns:
177
+ prediction (array) : the probability of the roi belonging to each class
178
+ '''
179
+ img_preprocessed = self.preprocessing(img_input)
180
+ prediction = self.predict(img_preprocessed)
181
+
182
+ return prediction