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
1721
1822
1923unet_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\n Average 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 )
0 commit comments