-
Notifications
You must be signed in to change notification settings - Fork 1
Expand file tree
/
Copy pathtrain.py
More file actions
739 lines (618 loc) · 35.2 KB
/
Copy pathtrain.py
File metadata and controls
739 lines (618 loc) · 35.2 KB
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
642
643
644
645
646
647
648
649
650
651
652
653
654
655
656
657
658
659
660
661
662
663
664
665
666
667
668
669
670
671
672
673
674
675
676
677
678
679
680
681
682
683
684
685
686
687
688
689
690
691
692
693
694
695
696
697
698
699
700
701
702
703
704
705
706
707
708
709
710
711
712
713
714
715
716
717
718
719
720
721
722
723
724
725
726
727
728
729
730
731
732
733
734
735
736
737
738
739
import os
import sys
import math
import re
import numpy as np
from datetime import datetime
import argparse
import torch
import torch.optim as optim
from torch.utils.data import DataLoader
from torch.utils.tensorboard import SummaryWriter
import torch.multiprocessing as mp
from torch.amp import GradScaler, autocast
import torch.distributed as dist
from torch.nn.parallel import DistributedDataParallel as DDP
from torch.utils.data.distributed import DistributedSampler
# Set multiprocessing start method to avoid CUDA context issues with DataLoader workers
try:
mp.set_start_method('spawn', force=True)
except RuntimeError:
pass # Already set
torch.backends.cudnn.benchmark = True
torch.set_float32_matmul_precision("high")
ROOT_DIR = os.path.dirname(os.path.abspath(__file__))
sys.path.append(os.path.join(ROOT_DIR, 'utils'))
sys.path.append(os.path.join(ROOT_DIR, 'models'))
sys.path.append(os.path.join(ROOT_DIR, 'dataset'))
from models.graspnet import GraspNet
from models.loss import get_loss
from dataset.graspnet_dataset import GraspNetDataset, spconv_collate_fn, load_grasp_labels, load_grasp_labels_lazy, SceneAwareSampler
from tqdm import tqdm
def freeze_for_stable_finetune(net, log_fn=print):
"""Freeze all parameters except the stable score head (conv_stable).
Used when fine-tuning a pretrained model to add stable score prediction.
"""
frozen_count = 0
trainable_count = 0
for name, param in net.named_parameters():
if 'conv_stable' in name:
param.requires_grad = True
trainable_count += param.numel()
else:
param.requires_grad = False
frozen_count += param.numel()
log_fn(f"Frozen {frozen_count:,} params, training {trainable_count:,} params (conv_stable only)")
parser = argparse.ArgumentParser()
parser.add_argument('--dataset_root', default=None, required=True)
parser.add_argument('--camera', default='kinect', help='Camera split [realsense/kinect]')
parser.add_argument('--train_split', default='train', help='Training split [train/train_all]')
parser.add_argument('--checkpoint_path', help='Model checkpoint path', default=None)
parser.add_argument('--model_name', type=str, default=None)
parser.add_argument('--log_dir', default='logs/log')
parser.add_argument('--num_point', type=int, default=15000, help='Point Number [default: 15000]')
parser.add_argument('--seed_feat_dim', default=512, type=int, help='Point wise feature dim')
parser.add_argument('--voxel_size', type=float, default=0.005, help='Voxel Size to process point clouds ')
parser.add_argument('--max_epoch', type=int, default=10, help='Epoch to run [default: 10]')
parser.add_argument('--batch_size', type=int, default=4, help='Batch Size during training [default: 4]')
parser.add_argument('--learning_rate', type=float, default=0.001, help='Initial learning rate [default: 0.001]')
parser.add_argument('--resume', action='store_true', default=False, help='Whether to resume from checkpoint')
parser.add_argument('--use_amp', action='store_true', default=False,
help='Use torch.cuda.amp for mixed-precision training')
parser.add_argument('--single_sample', action='store_true', default=False,
help='Overfit test: use only 1 training sample for 10 epochs')
parser.add_argument('--num_workers', type=int, default=0, help='Number of DataLoader workers [default: 0]')
parser.add_argument('--persistent_workers', action='store_true', default=False,
help='Keep workers alive between epochs (reduces memory overhead with num_workers>0)')
parser.add_argument('--lazy_grasp_labels', action='store_true', default=False,
help='Use lazy loading for grasp labels to reduce memory (useful with many workers)')
parser.add_argument('--weight_decay', type=float, default=0.0,
help='Weight decay for AdamW optimizer (recommended: 0.02-0.05 for transformers) [default: 0.0]')
parser.add_argument('--backbone', type=str, default='transformer', choices=['transformer', 'transformer_pretrained', 'sonata', 'pointnet2', 'resunet', 'resunet18', 'resunet_rgb', 'resunet18_rgb'],
help='Backbone architecture [default: transformer]. resunet=14D, resunet18=18D (more layers). sonata=self-supervised PTv3 (CVPR 2025). Use _rgb suffix for 6-channel RGB input.')
parser.add_argument('--grad_clip', type=float, default=0.0,
help='Gradient clipping max norm (recommended: 1.0-5.0 for transformers, 0 to disable) [default: 0.0]')
parser.add_argument('--ptv3_pretrained_path', type=str, default=None,
help='Path to PTv3 pretrained weights (.pth file). If not specified, uses models/pointcept/model_best.pth')
parser.add_argument('--enable_flash', action='store_true', default=False,
help='Enable flash attention in PTv3 backbone (requires flash_attn package)')
parser.add_argument('--accumulation_steps', type=int, default=1,
help='Gradient accumulation steps (simulate larger batch with batch_size=1) [default: 1]')
parser.add_argument('--backbone_lr_scale', type=float, default=None,
help='Learning rate multiplier for backbone (e.g., 0.1 for pretrained). Default: 0.1 for transformer_pretrained/sonata, 1.0 otherwise')
parser.add_argument('--layer_decay', type=float, default=None,
help='Layer-wise LR decay factor for pretrained backbones. Each encoder stage gets lr * layer_decay^(num_stages - stage). '
'Default: 0.65 for sonata/transformer_pretrained, 1.0 (disabled) otherwise')
parser.add_argument('--enable_stable_score', action='store_true', default=False,
help='Enable stable score prediction to penalize grasps that may cause tipping [default: False]')
parser.add_argument('--view_start', type=int, default=0,
help='Starting view index (inclusive) for each scene [default: 0]')
parser.add_argument('--view_end', type=int, default=256,
help='Ending view index (exclusive) for each scene [default: 256]')
parser.add_argument('--include_floor', action='store_true', default=False,
help='Include floor/table points in training (uses graspness_full/ labels, requires running generate_graspness_full.py first)')
parser.add_argument('--lambda_stable', type=float, default=10.0,
help='Weight for stable score loss term [default: 10.0]')
parser.add_argument('--graspness_threshold', type=float, default=0.1,
help='Threshold for graspness score filtering during forward pass [default: 0.1]')
parser.add_argument('--nsample', type=int, default=16,
help='Number of samples for cloud crop in GraspNet [default: 16]')
parser.add_argument('--cosine_lr', action='store_true', default=False,
help='Use cosine annealing LR schedule with warmup instead of exponential decay')
parser.add_argument('--warmup_epochs', type=int, default=2,
help='Number of warmup epochs for cosine LR schedule [default: 2]')
parser.add_argument('--finetune', action='store_true', default=False,
help='Fine-tune mode: load weights but reset epoch to 0 and skip optimizer state. Use with --checkpoint_path to fine-tune a vanilla model with stable score.')
parser.add_argument('--debug_feature_stats', action='store_true', default=False,
help='Print mean/std feature statistics at each backbone stage for first forward pass [default: False]')
parser.add_argument('--no_translation_aug', action='store_true', default=False,
help='Disable random translation in data augmentation (paper uses no translation) [default: False]')
parser.add_argument('--use_val', action='store_true', default=False,
help='Enable validation: use train_reduced (95 scenes) for training and val_train (5 scenes) for validation [default: False]')
# DDP arguments (set automatically by torchrun, but can be overridden)
parser.add_argument('--local_rank', type=int, default=-1,
help='Local rank for distributed training (set by torchrun)')
cfgs = parser.parse_args()
# Set layer_decay default (0.65 for pretrained transformers, 1.0 = disabled otherwise)
if cfgs.layer_decay is None:
cfgs.layer_decay = 0.65 if cfgs.backbone in ('transformer_pretrained', 'sonata') else 1.0
# Set backbone_lr_scale default: 1.0 when LLRD is active (it handles the scaling),
# 0.1 for pretrained without LLRD, 1.0 for non-pretrained
if cfgs.backbone_lr_scale is None:
if cfgs.layer_decay < 1.0:
cfgs.backbone_lr_scale = 1.0 # LLRD handles per-stage scaling
elif cfgs.backbone in ('transformer_pretrained', 'sonata'):
cfgs.backbone_lr_scale = 0.1 # flat scaling fallback
else:
cfgs.backbone_lr_scale = 1.0
# Auto-enable cosine LR with warmup for pretrained backbones (unless user explicitly set cosine_lr)
if cfgs.backbone in ('transformer_pretrained', 'sonata') and not cfgs.cosine_lr:
cfgs.cosine_lr = True
# Distributed Training Setup
def is_distributed():
"""Check if we're running in distributed mode."""
return dist.is_available() and dist.is_initialized()
def is_main_process():
"""Check if this is the main process (rank 0 or non-distributed)."""
if not is_distributed():
return True
return dist.get_rank() == 0
def setup_distributed():
"""Initialize distributed training if running with torchrun."""
# Check if we're running in distributed mode (torchrun sets these env vars)
if 'RANK' in os.environ and 'WORLD_SIZE' in os.environ:
rank = int(os.environ['RANK'])
world_size = int(os.environ['WORLD_SIZE'])
local_rank = int(os.environ.get('LOCAL_RANK', 0))
# Initialize the distributed backend
dist.init_process_group(
backend='nccl',
init_method='env://',
world_size=world_size,
rank=rank
)
# Set the device for this process
torch.cuda.set_device(local_rank)
return local_rank, rank, world_size
else:
# Not running distributed
return 0, 0, 1
def cleanup_distributed():
"""Clean up distributed training."""
if is_distributed():
dist.destroy_process_group()
# Initialize distributed training
LOCAL_RANK, _, WORLD_SIZE = setup_distributed()
EPOCH_CNT = 0
# Load checkpoint if resuming OR fine-tuning
CHECKPOINT_PATH = cfgs.checkpoint_path if (cfgs.resume or cfgs.finetune) else None
if not os.path.exists(cfgs.log_dir) and is_main_process():
os.makedirs(cfgs.log_dir)
# Only main process writes to log file
if is_main_process():
LOG_FOUT = open(os.path.join(cfgs.log_dir, 'log_train.txt'), 'a')
LOG_FOUT.write(str(cfgs) + '\n')
if is_distributed():
LOG_FOUT.write(f'Distributed training: world_size={WORLD_SIZE}, local_rank={LOCAL_RANK}\n')
else:
LOG_FOUT = None
def log_string(out_str):
if is_main_process():
LOG_FOUT.write(out_str + '\n')
LOG_FOUT.flush()
print(out_str)
# Init datasets and dataloaders
def my_worker_init_fn(worker_id):
np.random.seed(np.random.get_state()[1][0] + worker_id)
def create_dataloaders():
"""Create datasets and dataloaders. Each process creates its own."""
# Load grasp labels (use lazy loading if specified to save memory with multiple workers)
if cfgs.lazy_grasp_labels:
log_string("Using lazy loading for grasp labels (memory-efficient mode)")
grasp_labels = load_grasp_labels_lazy(cfgs.dataset_root)
else:
log_string("Loading all grasp labels into memory (~21GB)")
grasp_labels = load_grasp_labels(cfgs.dataset_root)
# Stable score settings (labels auto-computed by dataset if missing)
if cfgs.enable_stable_score:
log_string("Stable score prediction enabled (labels will be auto-computed if missing)")
use_rgb = cfgs.backbone.endswith('_rgb') or cfgs.backbone == 'transformer_pretrained'
# Determine training split (use train_reduced if validation is enabled)
actual_train_split = 'train_reduced' if cfgs.use_val and cfgs.train_split == 'train' else cfgs.train_split
train_dataset = GraspNetDataset(cfgs.dataset_root, grasp_labels=grasp_labels, camera=cfgs.camera, split=actual_train_split,
num_points=cfgs.num_point, voxel_size=cfgs.voxel_size,
remove_outlier=True, augment=True, load_label=True, use_rgb=use_rgb,
enable_stable_score=cfgs.enable_stable_score, view_start=cfgs.view_start, view_end=cfgs.view_end,
include_floor=cfgs.include_floor, augment_translation=not cfgs.no_translation_aug)
log_string(f'train dataset length: {len(train_dataset)} (split: {actual_train_split})')
# For overfitting test use only 1 sample repeated 256 times
if cfgs.single_sample:
from torch.utils.data import Subset
train_dataset = Subset(train_dataset, [0] * 256)
cfgs.max_epoch = 20
log_string('Single-sample overfitting test enabled: 256x repeated, max_epoch set to 20')
# Create samplers
# SceneAwareSampler groups samples by scene to maximize collision label cache hits
train_sampler = None
if is_distributed():
train_sampler = DistributedSampler(train_dataset, shuffle=True)
log_string(f'Using DistributedSampler with {WORLD_SIZE} processes')
elif cfgs.single_sample:
# For single-sample overfitting, don't use SceneAwareSampler (Subset doesn't have scenename)
train_sampler = None
log_string('Single-sample mode: using default sampler')
else:
# Use scene-aware sampling for cache-friendly data loading
train_sampler = SceneAwareSampler(train_dataset, shuffle=True)
log_string('Using SceneAwareSampler for cache-friendly collision label loading')
train_dataloader = DataLoader(train_dataset, batch_size=cfgs.batch_size,
shuffle=False, # Sampler handles shuffling
sampler=train_sampler,
num_workers=cfgs.num_workers, pin_memory=True,
persistent_workers=(cfgs.persistent_workers and cfgs.num_workers > 0),
worker_init_fn=my_worker_init_fn, collate_fn=spconv_collate_fn)
log_string('train dataloader length: ' + str(len(train_dataloader)))
# Create validation dataloader if enabled
val_dataloader = None
if cfgs.use_val:
val_dataset = GraspNetDataset(cfgs.dataset_root, grasp_labels=grasp_labels, camera=cfgs.camera, split='val_train',
num_points=cfgs.num_point, voxel_size=cfgs.voxel_size,
remove_outlier=True, augment=False, load_label=True, use_rgb=use_rgb,
enable_stable_score=cfgs.enable_stable_score, view_start=cfgs.view_start, view_end=cfgs.view_end,
include_floor=cfgs.include_floor, augment_translation=False)
log_string(f'val dataset length: {len(val_dataset)} (split: val_train, scenes 95-99)')
val_dataloader = DataLoader(val_dataset, batch_size=cfgs.batch_size,
shuffle=False, num_workers=cfgs.num_workers, pin_memory=True,
persistent_workers=(cfgs.persistent_workers and cfgs.num_workers > 0),
worker_init_fn=my_worker_init_fn, collate_fn=spconv_collate_fn)
log_string('val dataloader length: ' + str(len(val_dataloader)))
return train_dataloader, train_sampler, val_dataloader
def create_model_and_optimizer():
"""Create model, optimizer, and scaler. Each process creates its own."""
net = GraspNet(
seed_feat_dim=cfgs.seed_feat_dim,
is_training=True,
backbone=cfgs.backbone,
ptv3_pretrained_path=cfgs.ptv3_pretrained_path,
enable_flash=cfgs.enable_flash,
enable_stable_score=cfgs.enable_stable_score,
graspness_threshold=cfgs.graspness_threshold,
nsample=cfgs.nsample,
debug_feature_stats=cfgs.debug_feature_stats,
)
# Set device based on distributed or single-GPU mode
if is_distributed():
device = torch.device(f"cuda:{LOCAL_RANK}")
torch.cuda.set_device(device)
else:
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
net.to(device)
if cfgs.enable_stable_score:
log_string(f"Stable score prediction enabled (lambda_stable={cfgs.lambda_stable})")
head_lr = cfgs.learning_rate
backbone_base_lr = cfgs.learning_rate * cfgs.backbone_lr_scale
layer_decay = cfgs.layer_decay
weight_decay = cfgs.weight_decay if cfgs.weight_decay > 0 else 0.0
def _get_backbone_stage(name):
"""Return the encoder stage index for a backbone parameter, or -1 for embedding/fusion.
Sonata: backbone.encoder.enc.enc{0-4}.block* / backbone.encoder.embedding.*
PTv3: backbone.enc.enc{0-4}.block* / backbone.dec.dec{0-3}.* / backbone.embedding.*
"""
# Encoder stages (both Sonata and PTv3)
m = re.search(r'\.enc\.enc(\d+)\.', name)
if m:
return int(m.group(1))
# PTv3 decoder stages — treat at same depth as the encoder stage they mirror
# dec3 mirrors enc4, dec2 mirrors enc3, dec1 mirrors enc2, dec0 mirrors enc1
m = re.search(r'\.dec\.dec(\d+)\.', name)
if m:
return int(m.group(1)) + 1 # dec3→stage4, dec2→stage3, etc.
# Everything else (embedding, fusion_proj) → stage -1 (gets lowest LR)
return -1
# Determine number of encoder stages from the backbone
num_enc_stages = 5 # Both Sonata and PTv3 have enc0..enc4
# Build per-stage parameter groups for LLRD
# Stage LR: backbone_base_lr * layer_decay^(num_enc_stages - stage)
# Embedding (stage -1): backbone_base_lr * layer_decay^(num_enc_stages + 1) (lowest)
# fusion_proj: treated as head (randomly initialized, like output_proj)
head_decay_params = []
head_no_decay_params = []
# Dict: stage_idx -> {'decay': [...], 'no_decay': [...]}
backbone_stage_params = {}
for name, param in net.named_parameters():
if not param.requires_grad:
continue
# output_proj and fusion_proj are randomly initialized projection layers → head LR
is_backbone = name.startswith('backbone.') and 'output_proj' not in name and 'fusion_proj' not in name
is_norm = any(n in name.lower() for n in ['layernorm', 'layer_norm', 'batchnorm', 'batch_norm', '.bn.', '.norm.', '.norm1.', '.norm2.'])
is_bias = name.endswith('.bias')
no_decay = is_norm or is_bias
if is_backbone:
stage = _get_backbone_stage(name)
if stage not in backbone_stage_params:
backbone_stage_params[stage] = {'decay': [], 'no_decay': []}
if no_decay:
backbone_stage_params[stage]['no_decay'].append(param)
else:
backbone_stage_params[stage]['decay'].append(param)
else:
if no_decay:
head_no_decay_params.append(param)
else:
head_decay_params.append(param)
# Build optimizer param groups
param_groups = []
# Backbone groups with per-stage LR (LLRD)
for stage in sorted(backbone_stage_params.keys()):
if stage == -1:
# Embedding — deepest decay
stage_scale = layer_decay ** (num_enc_stages + 1)
stage_name = 'embedding'
else:
# Encoder/decoder stage — deeper stages get higher LR
stage_scale = layer_decay ** (num_enc_stages - stage)
stage_name = f'stage{stage}'
stage_lr = backbone_base_lr * stage_scale
if backbone_stage_params[stage]['decay']:
param_groups.append({
'params': backbone_stage_params[stage]['decay'],
'lr': stage_lr,
'weight_decay': weight_decay,
'name': f'backbone_{stage_name}_decay',
'backbone_stage': stage,
})
if backbone_stage_params[stage]['no_decay']:
param_groups.append({
'params': backbone_stage_params[stage]['no_decay'],
'lr': stage_lr,
'weight_decay': 0.0,
'name': f'backbone_{stage_name}_no_decay',
'backbone_stage': stage,
})
# Head groups
if head_decay_params:
param_groups.append({'params': head_decay_params, 'lr': head_lr, 'weight_decay': 0.0, 'name': 'head_decay'})
if head_no_decay_params:
param_groups.append({'params': head_no_decay_params, 'lr': head_lr, 'weight_decay': 0.0, 'name': 'head_no_decay'})
if cfgs.weight_decay > 0:
optimizer = optim.AdamW(param_groups)
log_string(f"Optimizer: AdamW with weight_decay={weight_decay}")
else:
optimizer = optim.Adam(param_groups)
log_string(f"Optimizer: Adam (no weight decay)")
log_string(f"Learning rates: backbone_base={backbone_base_lr:.6f} (scale={cfgs.backbone_lr_scale}), "
f"heads={head_lr:.6f}, layer_decay={layer_decay}")
log_string(f"LR schedule: {'cosine + warmup(' + str(cfgs.warmup_epochs) + ' epochs)' if cfgs.cosine_lr else 'exponential (0.95^epoch)'}")
for pg in param_groups:
log_string(f" {pg['name']}: {len(pg['params'])} params, lr={pg['lr']:.8f}")
# Initialize GradScaler for AMP (to prevent small gradients from underflowing to zero)
scaler = GradScaler(enabled=cfgs.use_amp and device.type == 'cuda')
if cfgs.use_amp:
log_string("Using Automatic Mixed Precision (AMP) training")
start_epoch = 0
if CHECKPOINT_PATH is not None and os.path.isfile(CHECKPOINT_PATH):
checkpoint = torch.load(CHECKPOINT_PATH, map_location=device)
# Handle both DDP and non-DDP checkpoints
state_dict = checkpoint['model_state_dict']
# Remove 'module.' prefix if loading from DDP checkpoint into non-DDP model
if not is_distributed() and any(k.startswith('module.') for k in state_dict.keys()):
state_dict = {k.replace('module.', ''): v for k, v in state_dict.items()}
net.load_state_dict(state_dict, strict=False)
# In finetune mode, skip optimizer state and reset epoch
if cfgs.finetune:
log_string("-> FINETUNE mode: loaded weights from %s, resetting epoch to 0" % CHECKPOINT_PATH)
if cfgs.enable_stable_score:
freeze_for_stable_finetune(net, log_string)
else:
optimizer.load_state_dict(checkpoint['optimizer_state_dict'])
if 'scaler_state_dict' in checkpoint and cfgs.use_amp:
scaler.load_state_dict(checkpoint['scaler_state_dict'])
start_epoch = checkpoint['epoch']
log_string("-> loaded checkpoint %s (epoch: %d)" % (CHECKPOINT_PATH, start_epoch))
# Wrap model in DDP if running distributed
if is_distributed():
# find_unused_parameters=True is needed because some parameters might not be used
# in every forward pass (e.g., stable_score_head when disabled)
net = DDP(net, device_ids=[LOCAL_RANK], output_device=LOCAL_RANK,
find_unused_parameters=True)
log_string(f"Model wrapped in DistributedDataParallel (device_ids=[{LOCAL_RANK}])")
return net, optimizer, scaler, start_epoch, device
def get_current_lr(epoch, base_lr):
"""Calculate LR based on schedule (exponential or cosine with warmup)."""
if cfgs.cosine_lr:
# Cosine annealing with linear warmup
if epoch < cfgs.warmup_epochs:
return base_lr * (epoch + 1) / cfgs.warmup_epochs
else:
progress = (epoch - cfgs.warmup_epochs) / max(1, cfgs.max_epoch - cfgs.warmup_epochs)
return base_lr * 0.5 * (1 + math.cos(math.pi * progress))
else:
# Exponential decay (original)
return base_lr * (0.95 ** epoch)
def adjust_learning_rate(optimizer, epoch):
"""Adjust LR for all param groups, maintaining their relative LLRD ratios."""
head_lr = get_current_lr(epoch, cfgs.learning_rate)
backbone_base_lr = get_current_lr(epoch, cfgs.learning_rate * cfgs.backbone_lr_scale)
num_enc_stages = 5
for param_group in optimizer.param_groups:
group_name = param_group.get('name', '')
if 'backbone' in group_name:
stage = param_group.get('backbone_stage', 0)
if stage == -1:
stage_scale = cfgs.layer_decay ** (num_enc_stages + 1)
else:
stage_scale = cfgs.layer_decay ** (num_enc_stages - stage)
param_group['lr'] = backbone_base_lr * stage_scale
else:
param_group['lr'] = head_lr
def train_one_epoch(net, optimizer, scaler, device, train_dataloader, train_writer):
stat_dict = {} # collect statistics
epoch_stat_dict = {} # collect epoch-level statistics
adjust_learning_rate(optimizer, EPOCH_CNT)
net.train()
batch_interval = 20
# Zero gradients at the start of each epoch
optimizer.zero_grad()
# Only show progress bar on main process to avoid duplicates
data_iter = tqdm(enumerate(train_dataloader), desc='Training',
disable=not is_main_process(), total=len(train_dataloader))
for batch_idx, batch_data_label in data_iter:
# Transfer to GPU with non_blocking=True for async copy (works with pin_memory=True)
for key in batch_data_label:
if 'list' in key:
for i in range(len(batch_data_label[key])):
for j in range(len(batch_data_label[key][i])):
batch_data_label[key][i][j] = batch_data_label[key][i][j].to(device, non_blocking=True)
else:
batch_data_label[key] = batch_data_label[key].to(device, non_blocking=True)
# Skip empty batches (0 voxels after sparse quantization)
if 'coors' in batch_data_label and batch_data_label['coors'].shape[0] == 0:
log_string(f'[Train] Skipping batch {batch_idx}: empty sparse tensor (0 voxels)')
continue
# Forward pass with autocast for mixed precision
try:
with autocast(enabled=cfgs.use_amp, device_type=device.type):
end_points = net(batch_data_label)
loss, end_points = get_loss(end_points,
enable_stable_score=cfgs.enable_stable_score,
lambda_stable=cfgs.lambda_stable)
# Scale loss for gradient accumulation
loss = loss / cfgs.accumulation_steps
except RuntimeError as e:
if "can't find suitable algorithm" in str(e) or "assert faild" in str(e):
log_string(f'[Train] Skipping batch {batch_idx}: spconv error - {e}')
continue
raise # Re-raise other RuntimeErrors
# Backward pass with gradient scaling
scaler.scale(loss).backward()
# Only step optimizer every accumulation_steps
if (batch_idx + 1) % cfgs.accumulation_steps == 0:
# Gradient clipping (important for transformer stability)
if cfgs.grad_clip > 0:
scaler.unscale_(optimizer)
torch.nn.utils.clip_grad_norm_(net.parameters(), cfgs.grad_clip)
scaler.step(optimizer)
scaler.update()
optimizer.zero_grad()
for key in end_points:
if 'loss' in key or 'acc' in key or 'prec' in key or 'recall' in key or 'count' in key:
if key not in stat_dict:
stat_dict[key] = 0
if key not in epoch_stat_dict:
epoch_stat_dict[key] = 0
loss_value = end_points[key].item()
stat_dict[key] += loss_value
epoch_stat_dict[key] += loss_value
if (batch_idx + 1) % batch_interval == 0:
log_string(' ----epoch: %03d ---- batch: %03d ----' % (EPOCH_CNT, batch_idx + 1))
for key in sorted(stat_dict.keys()):
if is_main_process() and train_writer is not None:
train_writer.add_scalar(key, stat_dict[key] / batch_interval,
(EPOCH_CNT * len(train_dataloader) + batch_idx) * cfgs.batch_size)
log_string('mean %s: %f' % (key, stat_dict[key] / batch_interval))
stat_dict[key] = 0
# Handle remaining gradients if num_batches not divisible by accumulation_steps
if len(train_dataloader) % cfgs.accumulation_steps != 0:
if cfgs.grad_clip > 0:
scaler.unscale_(optimizer)
torch.nn.utils.clip_grad_norm_(net.parameters(), cfgs.grad_clip)
scaler.step(optimizer)
scaler.update()
optimizer.zero_grad()
# Log epoch-level averages to TensorBoard (only main process)
num_batches = len(train_dataloader)
if is_main_process() and train_writer is not None:
for key in sorted(epoch_stat_dict.keys()):
avg_value = epoch_stat_dict[key] / num_batches
train_writer.add_scalar('epoch_' + key, avg_value, EPOCH_CNT)
# Flush to ensure data is written to disk
train_writer.flush()
# Return epoch average loss
return epoch_stat_dict['loss/overall_loss'] / num_batches if 'loss/overall_loss' in epoch_stat_dict else 0
@torch.no_grad()
def validate_one_epoch(net, device, val_dataloader, val_writer):
"""Run validation and return average loss."""
stat_dict = {}
net.eval()
data_iter = tqdm(enumerate(val_dataloader), desc='Validation',
disable=not is_main_process(), total=len(val_dataloader))
for batch_idx, batch_data_label in data_iter:
# Transfer to GPU
for key in batch_data_label:
if 'list' in key:
for i in range(len(batch_data_label[key])):
for j in range(len(batch_data_label[key][i])):
batch_data_label[key][i][j] = batch_data_label[key][i][j].to(device, non_blocking=True)
else:
batch_data_label[key] = batch_data_label[key].to(device, non_blocking=True)
# Skip empty batches (0 voxels after sparse quantization)
if 'coors' in batch_data_label and batch_data_label['coors'].shape[0] == 0:
log_string(f'[Val] Skipping batch {batch_idx}: empty sparse tensor (0 voxels)')
continue
# Forward pass with error handling for spconv edge cases
try:
with autocast(enabled=cfgs.use_amp, device_type=device.type):
end_points = net(batch_data_label)
loss, end_points = get_loss(end_points,
enable_stable_score=cfgs.enable_stable_score,
lambda_stable=cfgs.lambda_stable)
except RuntimeError as e:
if "can't find suitable algorithm" in str(e) or "assert faild" in str(e):
log_string(f'[Val] Skipping batch {batch_idx}: spconv error - {e}')
continue
raise # Re-raise other RuntimeErrors
# Accumulate statistics
for key in end_points:
if 'loss' in key or 'acc' in key or 'prec' in key or 'recall' in key or 'count' in key:
if key not in stat_dict:
stat_dict[key] = 0
stat_dict[key] += end_points[key].item()
# Compute averages and log
num_batches = len(val_dataloader)
log_string(' ---- Validation Results ----')
for key in sorted(stat_dict.keys()):
avg_value = stat_dict[key] / num_batches
log_string('val %s: %f' % (key, avg_value))
if is_main_process() and val_writer is not None:
val_writer.add_scalar('epoch_' + key, avg_value, EPOCH_CNT)
if is_main_process() and val_writer is not None:
val_writer.flush()
return stat_dict.get('loss/overall_loss', 0) / num_batches
def train(start_epoch, net, optimizer, scaler, device, train_dataloader,
train_writer, train_sampler=None, val_dataloader=None, val_writer=None):
global EPOCH_CNT
for epoch in range(start_epoch, cfgs.max_epoch):
# Set epoch on distributed sampler for proper shuffling
if train_sampler is not None:
train_sampler.set_epoch(epoch)
EPOCH_CNT = epoch
log_string('**** EPOCH %03d ****' % epoch)
log_string('Current learning rate: head=%.6f, backbone=%.6f' % (
get_current_lr(epoch, cfgs.learning_rate),
get_current_lr(epoch, cfgs.learning_rate * cfgs.backbone_lr_scale)
))
log_string(str(datetime.now()))
# Log learning rate to TensorBoard (only main process)
if is_main_process():
train_writer.add_scalar('learning_rate/head', get_current_lr(epoch, cfgs.learning_rate), epoch)
train_writer.add_scalar('learning_rate/backbone', get_current_lr(epoch, cfgs.learning_rate * cfgs.backbone_lr_scale), epoch)
# Reset numpy seed.
# REF: https://github.com/pytorch/pytorch/issues/5059
np.random.seed()
train_loss = train_one_epoch(net, optimizer, scaler, device, train_dataloader, train_writer)
# Run validation if enabled
if val_dataloader is not None:
val_loss = validate_one_epoch(net, device, val_dataloader, val_writer)
log_string(f'Epoch {epoch} - Train Loss: {train_loss:.6f}, Val Loss: {val_loss:.6f}')
# Save regular checkpoint (only from main process)
if is_main_process():
# Get the underlying model if wrapped in DDP
model_to_save = net.module if is_distributed() else net
save_dict = {'epoch': epoch + 1, 'optimizer_state_dict': optimizer.state_dict(),
'model_state_dict': model_to_save.state_dict(),
'scaler_state_dict': scaler.state_dict()}
torch.save(save_dict, os.path.join(cfgs.log_dir, cfgs.model_name + '_epoch' + str(epoch + 1).zfill(2) + '.tar'))
if __name__ == '__main__':
# Create dataloaders (each process creates its own for DDP)
TRAIN_DATALOADER, TRAIN_SAMPLER, VAL_DATALOADER = create_dataloaders()
# Create model, optimizer, and scaler (each process creates its own for DDP)
net, optimizer, scaler, start_epoch, device = create_model_and_optimizer()
# TensorBoard Visualizers (only main process writes to TensorBoard)
if is_main_process():
TRAIN_WRITER = SummaryWriter(os.path.join(cfgs.log_dir, 'train'))
VAL_WRITER = SummaryWriter(os.path.join(cfgs.log_dir, 'val')) if cfgs.use_val else None
else:
TRAIN_WRITER = None
VAL_WRITER = None
# Start training
try:
train(start_epoch, net, optimizer, scaler, device, TRAIN_DATALOADER,
TRAIN_WRITER, TRAIN_SAMPLER, VAL_DATALOADER, VAL_WRITER)
finally:
# Ensure TensorBoard writers are properly closed
if is_main_process():
TRAIN_WRITER.close()
if VAL_WRITER is not None:
VAL_WRITER.close()
# Clean up distributed training
cleanup_distributed()