File size: 10,992 Bytes
3549cf5
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Train AEF reconstruction, uniformity, consistency and text objectives."""

import json
import os
import sys
from pathlib import Path

import numpy as np
import torch
import yaml
from torch.nn.parallel import DistributedDataParallel
from torch.utils.data import DataLoader, Dataset, DistributedSampler


ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))
from model.alphaearthfoundations import AlphaEarthFoundations, compute_losses


class AEFDataset(Dataset):
    def __init__(self, path, config):
        self.data = np.load(path)
        self.input_sources = config["data"]["input_sources"]
        self.target_sources = config["data"]["target_sources"]
        format_version = str(self.data["format_version"])
        if format_version != config["data"]["format_version"]:
            raise ValueError(f"Expected format {config['data']['format_version']}, got {format_version}")
        for name, spec in self.input_sources.items():
            expected = (spec["timesteps"], spec["channels"], config["data"]["image_size"], config["data"]["image_size"])
            if self.data[name].shape[1:] != expected:
                raise ValueError(f"{name} has shape {self.data[name].shape[1:]}, expected {expected}")

    def __len__(self):
        return len(self.data["valid_period"])

    def __getitem__(self, index):
        item = {name: torch.from_numpy(self.data[name][index]).float() for name in self.input_sources}
        item.update({f"timestamps_{name}": torch.from_numpy(self.data[f"timestamps_{name}"][index]).long()
                     for name in self.input_sources})
        for name in self.input_sources:
            item[f"frame_available_{name}"] = torch.from_numpy(self.data[f"frame_available_{name}"][index])
            item[f"channel_available_{name}"] = torch.from_numpy(self.data[f"channel_available_{name}"][index])
            item[f"pixel_valid_{name}"] = torch.from_numpy(self.data[f"pixel_valid_{name}"][index]).float()
            item[f"geometry_{name}"] = torch.from_numpy(self.data[f"geometry_{name}"][index]).float()
        for name, spec in self.target_sources.items():
            if name in self.input_sources:
                continue
            values = torch.from_numpy(self.data[f"target_sequence_{name}"][index])
            item[f"target_sequence_{name}"] = values.long() if spec["type"] == "categorical" else values.float()
            item[f"target_timestamps_{name}"] = torch.from_numpy(self.data[f"target_timestamps_{name}"][index]).long()
            item[f"target_pixel_valid_{name}"] = torch.from_numpy(self.data[f"target_pixel_valid_{name}"][index]).float()
            item[f"target_geometry_{name}"] = torch.from_numpy(self.data[f"target_geometry_{name}"][index]).float()
            item[f"target_frame_available_{name}"] = torch.from_numpy(self.data[f"target_frame_available_{name}"][index])
        item["support_period"] = torch.from_numpy(self.data["support_period"][index]).long()
        item["valid_period"] = torch.from_numpy(self.data["valid_period"][index]).long()
        item["text_target"] = torch.from_numpy(self.data["text_target"][index]).float()
        return item


def device_from_config(config, local_rank=0):
    requested = config["runtime"]["device"]
    if requested == "auto":
        return torch.device("cuda", local_rank) if torch.cuda.is_available() else torch.device("cpu")
    return torch.device(requested)


def _select_indices(available):
    selected = []
    for row in available:
        candidates = torch.nonzero(row, as_tuple=False).flatten()
        choice = torch.randint(len(candidates), (), device=candidates.device)
        selected.append(candidates[choice])
    return torch.stack(selected).to(available.device)


def _gather(values, indices):
    return values[torch.arange(len(values), device=values.device), indices]


def unpack(batch, config, device, remove_input_targets=True):
    input_names, target_specs = config["data"]["input_sources"], config["data"]["target_sources"]
    sources, frame_available = {}, {}
    timestamps = {name: batch[f"timestamps_{name}"].to(device) for name in input_names}
    targets, masks, target_times, target_periods, geometry = {}, {}, {}, {}, {}
    for name in input_names:
        values = batch[name].to(device)
        available = batch[f"frame_available_{name}"].to(device).clone()
        channel_available = batch[f"channel_available_{name}"].to(device)
        values = values * channel_available[:, :, :, None, None]
        indices = _select_indices(available)
        targets[name] = _gather(values, indices)
        masks[name] = _gather(batch[f"pixel_valid_{name}"].to(device), indices) * _gather(channel_available, indices)[:, :, None, None]
        target_times[name] = _gather(timestamps[name], indices)
        target_periods[name] = torch.stack([
            target_times[name] - 5 * 86_400_000, target_times[name] + 5 * 86_400_000
        ], dim=1)
        geometry[name] = _gather(batch[f"geometry_{name}"].to(device), indices)
        if remove_input_targets:
            available[torch.arange(len(available), device=device), indices] = False
        sources[name] = values * available[:, :, None, None, None]
        frame_available[name] = available
    for name in target_specs:
        if name in input_names:
            continue
        available = batch[f"target_frame_available_{name}"].to(device)
        indices = _select_indices(available)
        targets[name] = _gather(batch[f"target_sequence_{name}"].to(device), indices)
        masks[name] = _gather(batch[f"target_pixel_valid_{name}"].to(device), indices)
        target_times[name] = _gather(batch[f"target_timestamps_{name}"].to(device), indices)
        target_periods[name] = torch.stack([
            target_times[name] - 5 * 86_400_000, target_times[name] + 5 * 86_400_000
        ], dim=1)
        geometry[name] = _gather(batch[f"target_geometry_{name}"].to(device), indices)
    return sources, timestamps, frame_available, targets, masks, target_times, target_periods, geometry


def perturb_sources(sources, frame_available, timestamps, support_period, config):
    perturbed, perturbed_available = {}, {}
    source_dropout, frame_dropout = config["train"]["source_dropout"], config["train"]["frame_dropout"]
    for name, values in sources.items():
        output, available = values + 0.01 * torch.randn_like(values), frame_available[name].clone()
        source_mask = torch.rand(values.shape[0], 1, 1, 1, 1, device=values.device) < source_dropout[name]
        strategy = torch.randint(3, (1,), device=values.device).item()
        if strategy == 0:
            dropped = torch.rand_like(available.float()) < frame_dropout[name]
        else:
            midpoint = support_period[:, :1] + (support_period[:, 1:] - support_period[:, :1]) // 2
            dropped = timestamps[name] >= midpoint if strategy == 1 else timestamps[name] < midpoint
        available = available & ~dropped & ~source_mask[:, 0, 0, 0]
        perturbed[name] = output * available[:, :, None, None, None]
        perturbed_available[name] = available
    return perturbed, perturbed_available


def main():
    config = yaml.safe_load((ROOT / "conf/config.yaml").read_text())
    torch.manual_seed(config["seed"])
    distributed = int(os.environ.get("WORLD_SIZE", "1")) > 1
    local_rank = int(os.environ.get("LOCAL_RANK", "0"))
    if distributed:
        torch.distributed.init_process_group("nccl" if torch.cuda.is_available() else "gloo")
    rank = torch.distributed.get_rank() if distributed else 0
    device = device_from_config(config, local_rank)
    if device.type == "cuda":
        torch.cuda.set_device(device)

    dataset = AEFDataset(ROOT / config["data"]["root"] / "train.npz", config)
    sampler = DistributedSampler(dataset, shuffle=True) if distributed else None
    loader = DataLoader(dataset, batch_size=config["train"]["batch_size"], sampler=sampler,
                        shuffle=sampler is None, num_workers=config["train"]["num_workers"])
    model = AlphaEarthFoundations(config["data"]["input_sources"], config["data"]["target_sources"], config["model"]).to(device)
    if distributed:
        model = DistributedDataParallel(model, device_ids=[local_rank] if device.type == "cuda" else None)
    optimizer = torch.optim.Adam(model.parameters(), lr=config["train"]["learning_rate"],
                                 weight_decay=config["train"]["weight_decay"])
    weights = {name: config["train"][f"{name}_weight"] for name in ("reconstruction", "uniformity", "consistency", "text")}
    history = []
    for epoch in range(config["train"]["epochs"]):
        if sampler:
            sampler.set_epoch(epoch)
        model.train()
        totals = {}
        for batch in loader:
            sources, timestamps, frame_available, targets, masks, target_times, target_periods, geometry = unpack(batch, config, device)
            valid_period = batch["valid_period"].to(device)
            teacher = model(sources, timestamps, valid_period, frame_available, target_times, geometry, target_periods)
            perturbed, perturbed_available = perturb_sources(
                sources, frame_available, timestamps, batch["support_period"].to(device), config
            )
            student = model(perturbed, timestamps, valid_period, perturbed_available)
            loss, components = compute_losses(teacher, student, targets, masks, batch["text_target"].to(device),
                                               config["data"]["target_sources"], weights)
            optimizer.zero_grad(set_to_none=True)
            loss.backward()
            torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
            optimizer.step()
            for name, value in components.items():
                totals[name] = totals.get(name, 0.0) + float(value.detach())
        metrics = {name: value / len(loader) for name, value in totals.items()}
        history.append({"epoch": epoch + 1, **metrics})
        if rank == 0:
            print(f"epoch={epoch + 1} total_loss={metrics['total']:.6f} reconstruction={metrics['reconstruction']:.6f}")

    if rank == 0:
        checkpoint = ROOT / config["paths"]["checkpoint"]
        metrics_path = ROOT / config["paths"]["training_metrics"]
        checkpoint.parent.mkdir(parents=True, exist_ok=True)
        metrics_path.parent.mkdir(parents=True, exist_ok=True)
        state = model.module.state_dict() if distributed else model.state_dict()
        torch.save({"model": state, "model_config": config["model"], "input_sources": config["data"]["input_sources"],
                    "target_sources": config["data"]["target_sources"], "format_version": config["data"]["format_version"]}, checkpoint)
        metrics_path.write_text(json.dumps({"history": history}, indent=2) + "\n")
        print(f"checkpoint={checkpoint.relative_to(ROOT)}")
    if distributed:
        torch.distributed.destroy_process_group()


if __name__ == "__main__":
    main()