| """One-batch phase-3 smoke test using the completed phase-2 checkpoint.""" |
|
|
| import os |
|
|
| |
| os.environ.setdefault("CUDA_VISIBLE_DEVICES", "2") |
|
|
| import torch |
| from torch.utils.data import DataLoader |
|
|
| from dataset.data_core import myDataset |
| from lib.config import config |
| from lib.utils import FullModel, get_model, get_optimizer |
|
|
|
|
| def main(): |
| config.defrost() |
| config.merge_from_file("lib/config/trufor_ph3.yaml") |
| config.GPUS = (0,) |
| config.WORKERS = 0 |
| config.DATASET.TRAIN = ["IMD", "CA", "CocoGlide"] |
| config.DATASET.VALID = ["IMD", "CA", "CocoGlide"] |
| smoke_batch = int(os.environ.get("PH3_SMOKE_BATCH", "1")) |
| config.TRAIN.BATCH_SIZE_PER_GPU = smoke_batch |
| config.TRAIN.NUM_SAMPLES = 1 |
| config.VALID.MAX_SIZE = 1024 |
| config.freeze() |
|
|
| assert torch.cuda.is_available(), "CUDA is unavailable" |
| print("CUDA:", torch.cuda.get_device_name(0), flush=True) |
|
|
| crop_size = (config.TRAIN.IMAGE_SIZE[1], config.TRAIN.IMAGE_SIZE[0]) |
| dataset = myDataset( |
| config, crop_size=crop_size, grid_crop=False, mode="train", aug=None |
| ) |
| print("Datasets:", dataset.get_info(), flush=True) |
|
|
| |
| for child in dataset.dataset_list: |
| rgb, label = child.get_img(0) |
| print(child.__class__.__name__, tuple(rgb.shape), tuple(label.shape), flush=True) |
|
|
| wrapped = torch.nn.DataParallel(get_model(config), device_ids=[0]).cuda() |
| model = FullModel(wrapped, config).cuda() |
|
|
| checkpoint_path = "weights/trufor_ph2/best.pth.tar" |
| checkpoint = torch.load(checkpoint_path, map_location="cpu") |
| incompatible = model.model.module.load_state_dict( |
| checkpoint["state_dict"], strict=False |
| ) |
| print("Phase-2 checkpoint epoch:", checkpoint.get("epoch", "unknown"), flush=True) |
| print("New phase-3 keys:", len(incompatible.missing_keys), flush=True) |
| print("Unexpected keys:", len(incompatible.unexpected_keys), flush=True) |
| del checkpoint |
|
|
| trainable = [ |
| name for name, parameter in model.model.module.named_parameters() |
| if parameter.requires_grad |
| ] |
| print("Trainable tensors:", len(trainable), flush=True) |
| print("Trainable examples:", trainable[:8], flush=True) |
|
|
| loader = DataLoader(dataset, batch_size=smoke_batch, shuffle=False, num_workers=0) |
| rgbs, labels = next(iter(loader)) |
| rgbs = rgbs.cuda(non_blocking=True) |
| labels = labels.long().cuda(non_blocking=True) |
|
|
| optimizer = get_optimizer(model, config) |
| model.train() |
| optimizer.zero_grad() |
| losses, outputs, confidence, detection = model(labels=labels, rgbs=rgbs) |
| loss = losses.mean() |
| assert torch.isfinite(loss), f"Non-finite phase-3 loss: {loss.item()}" |
| loss.backward() |
| optimizer.step() |
|
|
| print("Localization:", tuple(outputs.shape), flush=True) |
| print("Confidence:", tuple(confidence.shape), flush=True) |
| print("Detection:", tuple(detection.shape), flush=True) |
| print("Loss:", float(loss.detach().cpu()), flush=True) |
| print("PHASE3_SMOKE_TEST_OK", flush=True) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|