Skip to content

Commit d29c62c

Browse files
committed
Add distance weighting for map overlaps
1 parent 3984c38 commit d29c62c

10 files changed

Lines changed: 85 additions & 29 deletions

‎R/do_unet_map.R‎

Lines changed: 7 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -11,6 +11,9 @@
1111
#' @param which Which model(s) to use: `'all'`, `'full'`, or integer 1-5
1212
#' @param clip Optional clip extent
1313
#' @param write_probs If TRUE, write probability layers
14+
#' @param use_distance_weights If TRUE (default), weight patch contributions by
15+
#' distance to the nearest patch edge when averaging overlapping predictions.
16+
#' Reduces visible tile seams. Set FALSE for uniform averaging.
1417
#' @param mapid Map database id
1518
#' @param fitid Fit database id (for reference / logging)
1619
#' @param requirecuda If TRUE (default), abort immediately if CUDA is not available rather than
@@ -24,7 +27,8 @@
2427

2528
do_unet_map <- function(model, site, fit_result = 'fit01', result,
2629
which = 'all', clip = NULL,
27-
write_probs = FALSE, mapid = NULL, fitid = NULL,
30+
write_probs = FALSE, use_distance_weights = TRUE,
31+
mapid = NULL, fitid = NULL,
2832
requirecuda = TRUE, rep = NULL) {
2933

3034

@@ -130,7 +134,8 @@ do_unet_map <- function(model, site, fit_result = 'fit01', result,
130134
patches_dir = patches_dir,
131135
output_file = output_file,
132136
config = config,
133-
write_probs = write_probs
137+
write_probs = write_probs,
138+
use_distance_weights = use_distance_weights
134139
)
135140

136141

‎R/do_unet_prep_map.R‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -166,7 +166,7 @@ do_unet_prep_map <- function(model, clip = NULL) {
166166
model = model,
167167
n_patches = n_patches,
168168
patch_size = patch_size,
169-
overlap = MAP_OVERLAP,
169+
mapping_overlap = MAP_OVERLAP,
170170
stride = as.integer(stride),
171171
n_channels = n_channels,
172172
n_rows_rast = n_rows_rast,

‎R/map.R‎

Lines changed: 9 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -27,6 +27,10 @@
2727
#' model), or an integer CV fold number. Ignored for RF/AdaBoost models.
2828
#' @param write_probs For U-Net models: if TRUE, write per-class probability
2929
#' layers alongside the classification. Ignored for RF/AdaBoost models.
30+
#' @param use_distance_weights For U-Net models: if TRUE (default), weight
31+
#' patch contributions by distance to the nearest patch edge when averaging
32+
#' overlapping predictions, reducing visible tile seams. Set FALSE for
33+
#' uniform averaging. Ignored for RF/AdaBoost models.
3034
#' @param requirecuda If TRUE (default), abort immediately if CUDA is not available rather than
3135
#' silently falling back to CPU. Set to FALSE only for testing without a GPU.
3236
#' @param resources Slurm launch resources. See \link[slurmcollie]{launch}.
@@ -44,7 +48,8 @@
4448

4549

4650
map <- function(fit, site = NULL, clip = NULL, result = NULL,
47-
which = 'all', write_probs = FALSE, requirecuda = TRUE,
51+
which = 'all', write_probs = FALSE, use_distance_weights = TRUE,
52+
requirecuda = TRUE,
4853
resources = NULL, local = FALSE, trap = FALSE, comment = NULL) {
4954

5055

@@ -211,7 +216,9 @@ map <- function(fit, site = NULL, clip = NULL, result = NULL,
211216
launch('do_unet_map', reps = unet_model, repname = 'model',
212217
moreargs = list(site = site, fit_result = unet_fit_result,
213218
result = result, which = which, clip = clip,
214-
write_probs = write_probs, fitid = fitid,
219+
write_probs = write_probs,
220+
use_distance_weights = use_distance_weights,
221+
fitid = fitid,
215222
requirecuda = requirecuda,
216223
mapid = the$mdb$mapid[i]),
217224
finish = 'map_finish', callerid = the$mdb$mapid[i],

‎R/unet_assemble_map.R‎

Lines changed: 37 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -9,6 +9,10 @@
99
#' @param config Config list (from prep yaml, with `classes` and `site`)
1010
#' @param write_probs If TRUE, also write per-class probability layers as a
1111
#' multi-band GeoTIFF alongside the classification
12+
#' @param use_distance_weights If TRUE (default), weight pixel contributes
13+
#' by distance to the nearest patch edge during averaging. This reduces
14+
#' visible tile artifacts at patch boundaries. Set FALSE for uniform
15+
#' averaging (faster, but may show seams with low overlap).
1216
#' @importFrom terra rast ext crs values writeRaster
1317
#' @importFrom reticulate import
1418
#' @importFrom rasterPrep addColorTable makeNiceTif addVat
@@ -17,10 +21,10 @@
1721

1822

1923
unet_assemble_map <- function(patches_dir, output_file, config,
20-
write_probs = FALSE) {
24+
write_probs = FALSE, use_distance_weights = TRUE) {
2125

2226

23-
np <- import('numpy')
27+
np <- import('numpy')
2428

2529
site <- toupper(config$site)
2630

@@ -49,9 +53,24 @@ unet_assemble_map <- function(patches_dir, output_file, config,
4953

5054
# ----- Allocate accumulator matrices -----
5155
# Using plain matrices to avoid terra overhead during accumulation
52-
prob_accum <- array(0, dim = c(n_rows, n_cols, n_classes)) # summed probabilities
53-
count <- matrix(0L, nrow = n_rows, ncol = n_cols) # number of contributing patches
54-
nodata_accum <- matrix(0L, nrow = n_rows, ncol = n_cols) # nodata pixel count
56+
prob_accum <- array(0, dim = c(n_rows, n_cols, n_classes)) # summed probabilities
57+
count <- matrix(0, nrow = n_rows, ncol = n_cols) # sum of weights from contributing patches
58+
nodata_accum <- matrix(0L, nrow = n_rows, ncol = n_cols) # nodata pixel count
59+
60+
61+
# ----- Build distance-to-edge weight matrix -----
62+
# Weight = distance to nearest patch edge, normalized.
63+
# Center pixels get weight 1, edge pixels approach (but never reach) 0.
64+
# No zero weights ensures every pixel contributes something at raster boundaries.
65+
if(use_distance_weights) {
66+
if(patch_size %% 2 != 0)
67+
stop('patch_size must be even for distance weighting (got ', patch_size, ')')
68+
half <- patch_size / 2
69+
ramp <- c(seq_len(half), rev(seq_len(half))) / half # 1/half, 2/half, ..., 1, 1, ..., 2/half, 1/half
70+
edge_weight <- outer(ramp, ramp, pmin) # 2D pyramid: weight = distance to nearest edge
71+
} else {
72+
edge_weight <- matrix(1, nrow = patch_size, ncol = patch_size)
73+
}
5574

5675

5776
# ----- Accumulate -----
@@ -66,21 +85,22 @@ unet_assemble_map <- function(patches_dir, output_file, config,
6685
actual_w <- c1 - c0 + 1
6786

6887
nd_patch <- nodata[i, 1:actual_h, 1:actual_w] # nodata mask for this patch
88+
w_patch <- edge_weight[1:actual_h, 1:actual_w] * nd_patch # combined edge weight + nodata mask
6989

7090
for(k in seq_len(n_classes))
7191
prob_accum[r0:r1, c0:c1, k] <- prob_accum[r0:r1, c0:c1, k] +
72-
probs[i, k, 1:actual_h, 1:actual_w] * nd_patch # only accumulate valid pixels
92+
probs[i, k, 1:actual_h, 1:actual_w] * w_patch # weighted accumulation
7393

74-
count[r0:r1, c0:c1] <- count[r0:r1, c0:c1] + nd_patch
94+
count[r0:r1, c0:c1] <- count[r0:r1, c0:c1] + w_patch # sum of weights (now float, not integer)
7595
nodata_accum[r0:r1, c0:c1] <- nodata_accum[r0:r1, c0:c1] +
7696
as.integer(nd_patch == 0)
7797

7898
if(i %% 500 == 0)
7999
message(sprintf(' Processed %d / %d patches', i, n_patches))
80100
}
81-
101+
82102
rm(probs, nodata) # free memory
83-
103+
84104

85105
# ----- Average probabilities -----
86106
message('Averaging overlapping predictions...')
@@ -89,8 +109,8 @@ unet_assemble_map <- function(patches_dir, output_file, config,
89109

90110
cat('\n\nAverage probs peakRAM:\n')
91111
print(peakRAM({
92-
for(k in seq_len(n_classes))
93-
prob_accum[, , k] <- prob_accum[, , k] / count
112+
for(k in seq_len(n_classes))
113+
prob_accum[, , k] <- prob_accum[, , k] / count
94114
}))
95115

96116

@@ -113,26 +133,26 @@ unet_assemble_map <- function(patches_dir, output_file, config,
113133

114134
result_rast <- setValues(template, as.vector(t(pred_original))) # terra expects column-major, t() to match
115135

116-
136+
117137
# ----- Preliminary save -----
118138
dir.create(dirname(output_file), showWarnings = FALSE, recursive = TRUE)
119139
f0 <- file.path(dirname(output_file), paste0('zz_', basename(output_file), '_0'))
120140
writeRaster(result_rast, f0, overwrite = TRUE, datatype = 'INT1U')
121141

122-
142+
123143
# ----- Color table and VAT -----
124144
classes <- read_pars_table('classes')
125-
145+
126146
# Determine which class column to use (subclass, or reclassified e.g. ICS_V5)
127147
class_col <- if(!is.null(config$reclass) && nzchar(config$reclass)) config$reclass else 'subclass'
128148
name_col <- paste0(class_col, '_name')
129149
color_col <- paste0(class_col, '_color')
130-
150+
131151
# Build a deduplicated lookup table for the relevant class column
132152
class_lookup <- unique(classes[, c(class_col, name_col, color_col)])
133153
class_lookup <- class_lookup[!is.na(class_lookup[[class_col]]), ]
134154
names(class_lookup) <- c('subclass', 'name', 'color')
135-
155+
136156
# Build VAT from our predicted classes
137157
pred_classes <- sort(unique(as.vector(pred_original[!is.na(pred_original)])))
138158
vat <- data.frame(value = pred_classes, subclass = as.integer(pred_classes))
@@ -145,7 +165,7 @@ unet_assemble_map <- function(patches_dir, output_file, config,
145165
)
146166

147167
vrt_file <- addColorTable(f0, table = vat2)
148-
168+
149169
makeNiceTif(source = vrt_file, destination = output_file, overwrite = TRUE,
150170
overviewResample = 'nearest', stats = FALSE, vat = TRUE)
151171
addVat(output_file, attributes = vat)

‎man/compare.Rd‎

Lines changed: 5 additions & 3 deletions
Some generated files are not rendered by default. Learn more about customizing how changed files appear on GitHub.

‎man/describe_channels.Rd‎

Lines changed: 1 addition & 1 deletion
Some generated files are not rendered by default. Learn more about customizing how changed files appear on GitHub.

‎man/do_unet_map.Rd‎

Lines changed: 5 additions & 0 deletions
Some generated files are not rendered by default. Learn more about customizing how changed files appear on GitHub.

‎man/map.Rd‎

Lines changed: 6 additions & 0 deletions
Some generated files are not rendered by default. Learn more about customizing how changed files appear on GitHub.

‎man/parse_image_filename.Rd‎

Lines changed: 2 additions & 2 deletions
Some generated files are not rendered by default. Learn more about customizing how changed files appear on GitHub.

‎man/unet_assemble_map.Rd‎

Lines changed: 12 additions & 1 deletion
Some generated files are not rendered by default. Learn more about customizing how changed files appear on GitHub.

0 commit comments

Comments
 (0)