File size: 26,713 Bytes
c2b1b26
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
595
596
597
598
599
600
601
602
603
604
605
606
607
608
609
610
611
612
613
614
615
616
617
618
619
620
621
622
623
624
625
626
627
628
629
630
631
632
633
634
635
636
637
638
639
640
641
import json
import os
import sys
from PIL import Image
import numpy as np
import cv2
from osgeo import gdal, ogr, osr

# Increase PIL image size limit to handle large images
Image.MAX_IMAGE_PIXELS = None

# =======================================================
# I. 栅格处理工具函数 (基于您的代码,进行优化和整合)
# =======================================================

def create_geotiff_from_png(png_path, ref_tiff_path, out_tiff_path):
    """

    将拼接好的二值PNG图转换为具有地理参考信息的GeoTIFF。

    使用GDAL直接读取PNG,避免内存问题。

    """
    print(f"\n--- 🌐 步骤 II: 转换为 GeoTIFF ---")
    
    # 1. 打开原始 GeoTIFF 获取地理信息
    ds_ref = gdal.Open(ref_tiff_path, gdal.GA_ReadOnly)
    if ds_ref is None:
        raise RuntimeError(f"致命错误: 无法打开原始参考文件 {ref_tiff_path}")
    
    geo_transform = ds_ref.GetGeoTransform()
    projection = ds_ref.GetProjection()
    
    # 2. 使用GDAL直接打开PNG文件(避免PIL内存问题)
    print("⏳ 正在读取PNG文件...")
    try:
        # 使用GDAL打开PNG
        ds_png = gdal.Open(png_path, gdal.GA_ReadOnly)
        if ds_png is None:
            raise RuntimeError(f"无法使用GDAL打开PNG文件: {png_path}")
        
        # 获取PNG尺寸
        png_width = ds_png.RasterXSize
        png_height = ds_png.RasterYSize
        
        print(f"✅ PNG尺寸: {png_width} x {png_height}")
        
        # 读取PNG数据(分块读取以节省内存)
        png_band = ds_png.GetRasterBand(1)
        
    except Exception as e:
        # 如果GDAL无法打开PNG,尝试使用PIL(但分块处理)
        print(f"⚠️  GDAL无法打开PNG,尝试使用PIL分块读取...")
        try:
            img_png = Image.open(png_path)
            png_width, png_height = img_png.size
            
            # 分块读取(每次读取1000行)
            chunk_size = 1000
            data_chunks = []
            for y_start in range(0, png_height, chunk_size):
                y_end = min(y_start + chunk_size, png_height)
                box = (0, y_start, png_width, y_end)
                chunk = np.array(img_png.crop(box).convert('L'))
                data_chunks.append(chunk)
                if (y_start // chunk_size + 1) % 10 == 0:
                    print(f"   已读取 {y_end}/{png_height} 行...")
            
            # 合并数据块
            data = np.vstack(data_chunks)
            img_png.close()
            
        except Exception as e2:
            raise RuntimeError(f"读取PNG图像失败: {e2}")

    # 3. 创建新的 GeoTIFF 文件
    driver = gdal.GetDriverByName("GTiff")
    if os.path.exists(out_tiff_path):
        driver.Delete(out_tiff_path) # 确保覆盖旧文件
        
    print(f"⏳ 正在创建GeoTIFF文件...")
    ds_new = driver.Create(
        out_tiff_path,
        png_width,  # XSize (Width)
        png_height, # YSize (Height)
        1,          # Band Count
        gdal.GDT_Byte, # 确保使用 8位无符号整型存储二值数据 (0/1/255)
        options=['COMPRESS=DEFLATE', 'NUM_THREADS=ALL_CPUS', 'TILED=YES', 'BLOCKXSIZE=256', 'BLOCKYSIZE=256']
    )
    
    # 4. 设置地理信息和写入数据
    ds_new.SetGeoTransform(geo_transform)
    ds_new.SetProjection(projection)
    
    band = ds_new.GetRasterBand(1)
    
    # 分块写入数据(如果使用GDAL读取)
    if 'ds_png' in locals():
        print("⏳ 正在分块写入数据...")
        block_size = 1000  # 每次写入1000行
        for y_start in range(0, png_height, block_size):
            y_end = min(y_start + block_size, png_height)
            data_chunk = png_band.ReadAsArray(0, y_start, png_width, y_end - y_start)
            band.WriteArray(data_chunk, 0, y_start)
            if (y_start // block_size + 1) % 10 == 0:
                print(f"   已写入 {y_end}/{png_height} 行...")
        ds_png = None
    else:
        # 如果使用PIL读取,直接写入
        print("⏳ 正在写入数据...")
        band.WriteArray(data)
    
    band.SetNoDataValue(0) # 将背景 0 值设置为 NoData(可选)

    # 5. 清理资源
    ds_new = None
    ds_ref = None
    if 'data' in locals():
        del data  # 释放内存
    
    print(f"✅ GeoTIFF 转换成功,保存至: {out_tiff_path}")
    
    # 可选:立即添加颜色表(使用您的函数)
    add_transparent_color(out_tiff_path)

    return out_tiff_path

def add_transparent_color(raster_path):
    """把 0 值设成完全透明,1 值设成任意可见色 (适用于 GDT_Byte 或 GDT_UInt16)"""
    try:
        ds = gdal.Open(raster_path, gdal.GA_Update)
        if ds is None:
            print(f"警告: 无法以更新模式打开 {raster_path} 进行颜色表设置。")
            return

        band = ds.GetRasterBand(1)
        # 强制转换为 Byte 类型,如果不是的话,确保颜色表能生效
        if band.DataType != gdal.GDT_Byte:
            # 如果不是 Byte,先转换为 Byte(但这里假设 create_geotiff_from_png 已经处理了)
            print(f"警告: 栅格类型 {gdal.GetDataTypeName(band.DataType)} 可能不支持颜色表。")

        ct = gdal.ColorTable()
        ct.SetColorEntry(0, (0, 0, 0, 0))    # A=0  完全透明
        ct.SetColorEntry(1, (255, 0, 0, 180)) # A=180 半透明红色
        band.SetColorTable(ct)
        band.SetRasterColorInterpretation(gdal.GCI_PaletteIndex)
        ds = None
        print(f"✅ {os.path.basename(raster_path)} 已添加透明颜色表。")
    except Exception as e:
        print(f"❌ 颜色表设置失败: {e}")

def raster2polygon(in_raster_path, out_shp_path,

                   field_name='DN',

                   connected_8=True,

                   sieve_size=0,

                   target_value=255,

                   contour_approx_epsilon=None,

                   smooth_gaussian=True,

                   gaussian_kernel_size=5,

                   use_spline=False,

                   spline_points=100,

                   min_area=0.0,

                   check_topology=True):
    """

    使用OpenCV轮廓检测和轮廓近似进行矢量化,获得平滑边界(减少锐角)

    

    Args:

        in_raster_path: 输入栅格路径

        out_shp_path: 输出矢量路径

        field_name: 字段名

        connected_8: 是否使用8连通(保留参数以兼容旧代码,但新方法使用OpenCV轮廓检测)

        sieve_size: 碎斑过滤阈值(像素数)

        target_value: 目标像素值

        contour_approx_epsilon: 轮廓近似精度(像素),None表示自动计算

                                值越小越精确但顶点越多,值越大越平滑但可能丢失细节

                                建议值:0.5-2.0像素(减小以获得更平滑的边界)

        smooth_gaussian: 是否在矢量化前对栅格进行高斯模糊平滑(默认True)

        gaussian_kernel_size: 高斯模糊核大小(奇数,建议3-7,默认5)

        use_spline: 是否使用样条插值进一步平滑边界(默认False,会增加计算时间)

        spline_points: 样条插值的点数(仅在use_spline=True时有效)

        min_area: 最小面积阈值(平方米),小于此值的多边形将被过滤(默认0.0不过滤)

        check_topology: 是否进行拓扑检查和修复(默认True)

    """
    print(f"\n--- 🗺️ 步骤 III: 栅格转矢量(轮廓平滑方法)---")
    
    ds = gdal.Open(in_raster_path, gdal.GA_ReadOnly)
    if ds is None:
        raise RuntimeError(f'无法打开输入栅格:{in_raster_path}')

    band = ds.GetRasterBand(1)
    data = band.ReadAsArray()
    
    # 获取地理变换参数
    geo_transform = ds.GetGeoTransform()
    projection = ds.GetProjection()
    
    # 统计
    target_pixels = np.sum(data == target_value)
    total_pixels = data.size
    print(f"目标像素(值={target_value}): {target_pixels:,} / {total_pixels:,}")
    
    if target_pixels == 0:
        print("⚠️ 没有找到目标像素,跳过矢量化")
        ds = None
        return None
    
    # 创建掩膜
    mask_data = np.where(data == target_value, 255, 0).astype(np.uint8)
    
    # 碎斑过滤(使用OpenCV)
    if sieve_size > 0:
        print(f"⏳ 碎斑过滤 (阈值: {sieve_size} 像素)...")
        # 使用形态学开运算去除小碎斑
        kernel_size = max(3, sieve_size // 2)
        if kernel_size % 2 == 0:
            kernel_size += 1
        kernel = cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (kernel_size, kernel_size))
        mask_data = cv2.morphologyEx(mask_data, cv2.MORPH_OPEN, kernel)
    
    # 高斯模糊平滑(减少锐角,使边界更平滑)
    if smooth_gaussian:
        print(f"⏳ 高斯模糊平滑 (核大小: {gaussian_kernel_size})...")
        # 确保核大小为奇数
        if gaussian_kernel_size % 2 == 0:
            gaussian_kernel_size += 1
        mask_data = cv2.GaussianBlur(mask_data, (gaussian_kernel_size, gaussian_kernel_size), 0)
        # 重新二值化(高斯模糊后值会变化)
        _, mask_data = cv2.threshold(mask_data, 127, 255, cv2.THRESH_BINARY)
    
    # 创建输出目录
    out_dir = os.path.dirname(out_shp_path)
    if out_dir and not os.path.exists(out_dir):
        os.makedirs(out_dir, exist_ok=True)
    
    # 删除已存在的文件
    if os.path.exists(out_shp_path):
        for ext in ['.shp', '.shx', '.dbf', '.prj', '.qpj']:
            try:
                os.remove(out_shp_path.replace('.shp', ext))
            except:
                pass
    
    # 创建矢量文件
    driver = ogr.GetDriverByName('ESRI Shapefile')
    out_ds = driver.CreateDataSource(out_shp_path)
    
    srs = osr.SpatialReference()
    srs.ImportFromWkt(projection)
    
    out_layer = out_ds.CreateLayer(
        name=os.path.splitext(os.path.basename(out_shp_path))[0],
        geom_type=ogr.wkbPolygon,
        srs=srs
    )
    
    fd = ogr.FieldDefn(field_name, ogr.OFTInteger)
    out_layer.CreateField(fd)
    
    # 使用OpenCV轮廓检测
    print("⏳ 正在检测轮廓...")
    # 使用CHAIN_APPROX_NONE获取所有点,然后进行平滑处理(减少锐角)
    contours, hierarchy = cv2.findContours(
        mask_data, 
        cv2.RETR_CCOMP,  # 检测所有轮廓,包括孔洞
        cv2.CHAIN_APPROX_NONE  # 获取所有点,便于后续平滑处理
    )
    
    print(f"   检测到 {len(contours)} 个轮廓")
    
    # 计算轮廓近似精度(如果未指定)
    if contour_approx_epsilon is None:
        # 根据图像分辨率自动计算
        # 为了获得更平滑的边界,使用更小的epsilon值(增加顶点数)
        pixel_size_x = abs(geo_transform[1])
        pixel_size_y = abs(geo_transform[5])
        avg_pixel_size = (pixel_size_x + pixel_size_y) / 2
        
        # 转换为像素单位(假设地理坐标单位是度)
        # 对于约16米分辨率,约0.00015度,对应约1-2像素
        # 使用更小的值以获得更平滑的边界(减少锐角)
        if avg_pixel_size < 0.001:  # 地理坐标系(度)
            contour_approx_epsilon = 1.0  # 从2.0减小到1.0,增加顶点数
        else:  # 投影坐标系(米)
            contour_approx_epsilon = max(0.5, avg_pixel_size / 16.0)  # 从8.0改为16.0,更精细
    
    print(f"   轮廓近似精度: {contour_approx_epsilon:.2f} 像素(值越小边界越平滑)")
    
    # 辅助函数:样条插值平滑轮廓
    def smooth_contour_with_spline(contour, num_points=100):
        """

        使用样条插值平滑轮廓,减少锐角

        

        Args:

            contour: OpenCV轮廓点

            num_points: 插值后的点数

        

        Returns:

            平滑后的轮廓点

        """
        if len(contour) < 4:
            return contour
        
        # 提取x和y坐标
        points = contour.reshape(-1, 2)
        x = points[:, 0].astype(np.float32)
        y = points[:, 1].astype(np.float32)
        
        # 闭合轮廓(添加第一个点到末尾)
        x = np.append(x, x[0])
        y = np.append(y, y[0])
        
        # 计算累积距离作为参数
        distances = np.zeros(len(x))
        for i in range(1, len(x)):
            dx = x[i] - x[i-1]
            dy = y[i] - y[i-1]
            distances[i] = distances[i-1] + np.sqrt(dx*dx + dy*dy)
        
        # 归一化参数到[0, 1]
        if distances[-1] > 0:
            t = distances / distances[-1]
        else:
            return contour
        
        # 生成新的参数点
        t_new = np.linspace(0, 1, num_points)
        
        # 样条插值
        try:
            from scipy.interpolate import interp1d
            # 使用三次样条插值
            fx = interp1d(t, x, kind='cubic', bounds_error=False, fill_value='extrapolate')
            fy = interp1d(t, y, kind='cubic', bounds_error=False, fill_value='extrapolate')
            
            x_new = fx(t_new)
            y_new = fy(t_new)
            
            # 转换为OpenCV轮廓格式
            smoothed = np.array([[int(x_new[i]), int(y_new[i])] for i in range(len(x_new))], dtype=np.int32)
            return smoothed.reshape(-1, 1, 2)
        except ImportError:
            # 如果没有scipy,使用简单的线性插值
            print("   警告: 未安装scipy,使用线性插值代替样条插值")
            fx = np.interp(t_new, t, x)
            fy = np.interp(t_new, t, y)
            smoothed = np.array([[int(fx[i]), int(fy[i])] for i in range(len(fx))], dtype=np.int32)
            return smoothed.reshape(-1, 1, 2)
    
    # 辅助函数:将轮廓转换为OGR环
    def contour_to_ring(contour, geo_transform):
        """将OpenCV轮廓转换为OGR线性环"""
        ring = ogr.Geometry(ogr.wkbLinearRing)
        for point in contour:
            x_pixel = point[0][0]
            y_pixel = point[0][1]
            
            # 转换为地理坐标
            x_geo = geo_transform[0] + x_pixel * geo_transform[1] + y_pixel * geo_transform[2]
            y_geo = geo_transform[3] + x_pixel * geo_transform[4] + y_pixel * geo_transform[5]
            
            ring.AddPoint(x_geo, y_geo)
        
        # 闭合环
        if ring.GetPointCount() > 0:
            first_point = ring.GetPoint(0)
            ring.AddPoint(first_point[0], first_point[1])
        
        return ring
    
    # 处理轮廓:先处理外环,再处理对应的内环(孔洞)
    feature_count = 0
    processed_indices = set()
    
    for i, contour in enumerate(contours):
        if i in processed_indices or len(contour) < 3:
            continue
        
        # 检查是否是外环(hierarchy[i][3] == -1)
        parent_idx = hierarchy[0][i][3]
        if parent_idx != -1:  # 这是内环(孔洞),跳过,稍后处理
            continue
        
        # 轮廓处理:先进行样条插值平滑(如果启用),再进行Douglas-Peucker近似
        processed_contour = contour
        
        # 样条插值平滑(减少锐角)
        if use_spline and len(contour) >= 4:
            processed_contour = smooth_contour_with_spline(contour, spline_points)
        
        # 轮廓近似(Douglas-Peucker算法)- 使用较小的epsilon以获得更平滑的边界
        approx = cv2.approxPolyDP(processed_contour, contour_approx_epsilon, closed=True)
        
        if len(approx) < 3:
            continue
        
        # 创建外环
        exterior_ring = contour_to_ring(approx, geo_transform)
        
        # 创建多边形
        poly = ogr.Geometry(ogr.wkbPolygon)
        poly.AddGeometry(exterior_ring)
        
        # 查找并添加内环(孔洞)
        # 遍历所有轮廓,找到父轮廓是当前轮廓的内环
        child_idx = hierarchy[0][i][2]  # 第一个子轮廓索引
        while child_idx != -1:
            if child_idx < len(contours):
                child_contour = contours[child_idx]
                if len(child_contour) >= 3:
                    # 对内环也进行平滑处理
                    processed_child = child_contour
                    if use_spline and len(child_contour) >= 4:
                        processed_child = smooth_contour_with_spline(child_contour, spline_points)
                    
                    # 近似内环
                    child_approx = cv2.approxPolyDP(processed_child, contour_approx_epsilon, closed=True)
                    if len(child_approx) >= 3:
                        interior_ring = contour_to_ring(child_approx, geo_transform)
                        poly.AddGeometry(interior_ring)
                    processed_indices.add(child_idx)
                # 移动到下一个兄弟轮廓
                child_idx = hierarchy[0][child_idx][0]
            else:
                break
        
        # 拓扑检查和修复
        if check_topology:
            # 检查几何有效性
            if not poly.IsValid():
                # 尝试修复无效几何(Buffer(0)可以修复一些拓扑错误)
                try:
                    poly_fixed = poly.Buffer(0)
                    if poly_fixed.IsValid():
                        poly = poly_fixed
                        print(f"   ✓ 修复无效几何(轮廓 {i})")
                    else:
                        print(f"   ✗ 无法修复无效几何(轮廓 {i}),跳过")
                        processed_indices.add(i)
                        continue
                except Exception as e:
                    print(f"   ✗ 修复几何失败(轮廓 {i}): {e},跳过")
                    processed_indices.add(i)
                    continue
            
            # 检查面积(过滤太小的多边形)
            area = poly.GetArea()
            if area <= 0:
                print(f"   ✗ 跳过零面积几何(轮廓 {i})")
                processed_indices.add(i)
                continue
            
            # 如果设置了最小面积阈值,过滤太小的多边形
            if min_area > 0 and area < min_area:
                print(f"   ✗ 跳过面积过小的几何(轮廓 {i},面积: {area:.2f} 平方米)")
                processed_indices.add(i)
                continue
        
        # 创建要素
        feature = ogr.Feature(out_layer.GetLayerDefn())
        feature.SetGeometry(poly)
        feature.SetField(field_name, target_value)
        out_layer.CreateFeature(feature)
        feature = None
        feature_count += 1
        processed_indices.add(i)
    
    # 清理
    out_ds = None
    ds = None
    
    print(f"✅ 矢量化完成,共 {feature_count} 个面要素")
    print(f"   输出: {out_shp_path}")
    
    return out_shp_path

# =======================================================
# II. 拼接主函数 (基于上一次修正)
# =======================================================

def stitch_binary_predictions_patch_final(output_dir, json_filename, output_filename="stitched_prediction.png"):
    """

    拼接切片为 PNG 文件,确保与 'patch00000_prediction.png' 命名规则匹配。

    (代码与上一个回答的最终版本一致,这里仅作为整合)

    """
    json_filepath = os.path.join(output_dir, json_filename)
    print(f"--- 🚀 步骤 I: 拼接二值图 ---")
    
    try:
        with open(json_filepath, 'r') as f:
            data = json.load(f)
    except Exception as e:
        raise RuntimeError(f"❌ 错误: 读取或解析JSON文件出错: {e}")

    tile_results = data.get('tile_results', [])
    if not tile_results:
        print("⚠️ 警告: JSON文件中没有找到 'tile_results' 数据。")
        return None

    def get_tile_filename(tile_id):
        return f"patch{tile_id:05d}_prediction.png"

    first_tile_id = tile_results[0]['tile_id']
    first_tile_path = os.path.join(output_dir, get_tile_filename(first_tile_id))
    
    try:
        with Image.open(first_tile_path) as img:
            tile_width, tile_height = img.size
    except Exception as e:
        raise RuntimeError(f"❌ 致命错误: 无法读取第一个切片 {first_tile_path}。请确认文件名。")

    max_x = max(item['x'] for item in tile_results)
    max_y = max(item['y'] for item in tile_results)
    stitched_width = max_x + tile_width
    stitched_height = max_y + tile_height
    
    print(f"✅ 切片尺寸 (W x H): {tile_width} x {tile_height} | 预计大图尺寸: {stitched_width} x {stitched_height}")
    
    # 检查内存需求(估算)
    estimated_memory_mb = (stitched_width * stitched_height * 1) / (1024 * 1024)  # 1字节每像素
    print(f"📊 预计内存需求: {estimated_memory_mb:.2f} MB")
    
    if estimated_memory_mb > 10000:  # 超过10GB
        print("⚠️  警告: 图像很大,可能需要大量内存")
    
    try:
        stitched_image = Image.new('L', (stitched_width, stitched_height))
    except MemoryError:
        error_msg = (f"❌ 内存不足: 无法创建 {stitched_width}x{stitched_height} 的图像。\n"
                    f"   预计需要约 {estimated_memory_mb:.2f} MB 内存。\n"
                    f"   建议: 1) 关闭其他程序 2) 使用更大的内存 3) 分块处理")
        raise RuntimeError(error_msg)
    except Exception as e:
        raise RuntimeError(f"❌ 创建拼接图像失败: {e}")
    
    print(f"⏳ 开始拼接 {len(tile_results)} 个切片...")
    processed_count = 0
    
    for tile_info in tile_results:
        tile_id = tile_info['tile_id']
        x_offset = tile_info['x']
        y_offset = tile_info['y']
        tile_path = os.path.join(output_dir, get_tile_filename(tile_id))

        try:
            with Image.open(tile_path) as tile_img:
                if tile_img.mode != 'L':
                    tile_img = tile_img.convert('L')
                stitched_image.paste(tile_img, (x_offset, y_offset))
                processed_count += 1
                if processed_count % 1000 == 0:
                    print(f"   已处理 {processed_count}/{len(tile_results)} 个切片...")
        except FileNotFoundError:
            print(f"⚠️ 警告: 找不到切片图像文件,跳过: {tile_path}")
        except Exception as e:
            print(f"❌ 处理切片 {tile_id} 时出错,跳过: {e}")
    
    print(f"✅ 完成拼接,共处理 {processed_count} 个切片")

    final_output_path = os.path.join(output_dir, output_filename)
    
    # 直接保存,避免转换为numpy数组(节省内存)
    print(f"⏳ 正在保存拼接结果到 {output_filename}...")
    try:
        # 直接保存PIL Image,不需要转换为numpy
        stitched_image.save(final_output_path, format='PNG', compress_level=1)
        print(f"✅ PNG 拼接完成,保存至: {final_output_path}")
    except MemoryError:
        # 如果直接保存也失败,尝试使用更低的压缩级别或分块保存
        print("⚠️  直接保存失败,尝试使用分块保存...")
        try:
            # 尝试使用更低的压缩级别
            stitched_image.save(final_output_path, format='PNG', compress_level=9, optimize=True)
            print(f"✅ PNG 拼接完成(使用压缩),保存至: {final_output_path}")
        except Exception as e:
            raise RuntimeError(f"❌ 保存失败: {e}\n"
                             f"   图像太大 ({stitched_width}x{stitched_height}),内存不足。\n"
                             f"   建议: 1) 关闭其他程序 2) 增加虚拟内存 3) 使用分块处理")
    except Exception as e:
        raise RuntimeError(f"❌ 保存PNG文件失败: {e}")
    
    # 清理内存
    stitched_image = None
    
    return final_output_path

# =======================================================
# III. 主流程调用
# =======================================================
if __name__ == '__main__':
    # ------------------- 配置文件路径 -------------------
    # 存放切片和 JSON 的目录
    PREDICTION_OUTPUT_DIR = "outputs/test_large" 
    # JSON 文件名
    JSON_FILENAME = "GF6_PMS_E121.1_N33.6_20250604_L1A1420584616-MUX_fuse_tile_results.json" 
    
    # 原始影像路径 (用于获取地理参考信息,您在 JSON 中提到过文件名)
    ORIGINAL_TIFF_PATH = r"datu/test/GF6_PMS_E121.1_N33.6_20250604_L1A1420584616-MUX_fuse.tif" 
    
    # 输出文件名定义
    STITCHED_PNG_FILENAME = "stitched_prediction.png"
    OUTPUT_GEOTIFF_FILENAME = "final_binary_prediction.tif"
    OUTPUT_SHP_FILENAME = "output_vector_test/sargassum_polygon.shp"

    # 碎斑过滤设置
    SIEVE_PIXELS = 0
    # ----------------------------------------------------
    
    try:
        # 1. 拼接切片 (生成 PNG)
        stitched_png_path = stitch_binary_predictions_patch_final(
            PREDICTION_OUTPUT_DIR, 
            JSON_FILENAME, 
            STITCHED_PNG_FILENAME
        )
        if not stitched_png_path:
            sys.exit(1)

        # 2. 转换为 GeoTIFF (添加地理参考)
        stitched_tiff_path = create_geotiff_from_png(
            stitched_png_path,
            ORIGINAL_TIFF_PATH,
            os.path.join(PREDICTION_OUTPUT_DIR, OUTPUT_GEOTIFF_FILENAME)
        )
        
        # 3. 栅格矢量化
        raster2polygon(
            in_raster_path=stitched_tiff_path,
            out_shp_path=os.path.join(PREDICTION_OUTPUT_DIR, OUTPUT_SHP_FILENAME),
            sieve_size=SIEVE_PIXELS,
            connected_8=True
        )

    except Exception as e:
        import traceback
        print(f"\n❌ 致命错误发生: {e}")
        print("\n详细错误信息:")
        traceback.print_exc()
        print("\n可能的原因:")
        print("1. 文件路径不正确或文件不存在")
        print("2. 内存不足(图像太大)")
        print("3. 文件权限问题")
        print("4. GDAL库配置问题")
        sys.exit(1)