Sherlockhu commited on
Commit
fa697b8
·
verified ·
1 Parent(s): fdfc898

Add training script

Browse files
Files changed (1) hide show
  1. train_drone_vla.py +189 -0
train_drone_vla.py ADDED
@@ -0,0 +1,189 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ """
3
+ 🚁 Drone VLA Training: Qwen3.5-0.8B + LoRA + Action Head → [Vx, Vy, Vz, Yaw_rate]
4
+
5
+ Usage:
6
+ python train_drone_vla.py --dataset_repo YOUR_USER/drone-nav-data --hub_model_id YOUR_USER/drone-vla --num_epochs 10
7
+ """
8
+ import os, sys, json, math, argparse, numpy as np
9
+ from pathlib import Path
10
+ import torch, torch.nn as nn, torch.nn.functional as F
11
+ from torch.utils.data import DataLoader, Dataset as TorchDataset
12
+ from transformers import AutoProcessor, AutoModelForImageTextToText
13
+ from peft import LoraConfig, get_peft_model, TaskType
14
+ from datasets import load_from_disk, load_dataset
15
+ from PIL import Image
16
+ try:
17
+ import trackio; HAS_TRACKIO = True
18
+ except ImportError:
19
+ HAS_TRACKIO = False
20
+
21
+ class ActionHead(nn.Module):
22
+ def __init__(self, hidden_size=1024, action_dim=4, hidden_dim=256):
23
+ super().__init__()
24
+ self.net = nn.Sequential(
25
+ nn.Linear(hidden_size, hidden_dim), nn.GELU(), nn.Dropout(0.1),
26
+ nn.Linear(hidden_dim, hidden_dim), nn.GELU(), nn.Dropout(0.1),
27
+ nn.Linear(hidden_dim, action_dim), nn.Tanh())
28
+ def forward(self, x): return self.net(x)
29
+
30
+ class DroneVLA(nn.Module):
31
+ def __init__(self, model_name="Qwen/Qwen3.5-0.8B", action_dim=4, lora_r=16, lora_alpha=32, freeze_vision=True):
32
+ super().__init__()
33
+ print(f"Loading {model_name}...")
34
+ self.backbone = AutoModelForImageTextToText.from_pretrained(model_name, torch_dtype=torch.bfloat16, trust_remote_code=True)
35
+ hs = getattr(self.backbone.config, 'hidden_size', None) or getattr(self.backbone.config, 'text_config', type('',(),{'hidden_size':1024})).hidden_size
36
+ print(f"Hidden size: {hs}")
37
+ exclude = ["lm_head"]
38
+ if freeze_vision:
39
+ for name, _ in self.backbone.named_modules():
40
+ if any(v in name.lower() for v in ['visual','vision','vit','patch_embed']) and name.count('.') <= 1:
41
+ exclude.append(name)
42
+ exclude = list(set(exclude)) or ["visual", "lm_head"]
43
+ lora_config = LoraConfig(r=lora_r, lora_alpha=lora_alpha, lora_dropout=0.05, bias="none",
44
+ task_type=TaskType.CAUSAL_LM, target_modules="all-linear", exclude_modules=exclude)
45
+ self.backbone = get_peft_model(self.backbone, lora_config)
46
+ self.backbone.print_trainable_parameters()
47
+ self.action_head = ActionHead(hidden_size=hs, action_dim=action_dim)
48
+ self.hidden_size = hs
49
+
50
+ def forward(self, input_ids, attention_mask, pixel_values=None, image_grid_thw=None, action_targets=None, **kw):
51
+ bkw = {'input_ids': input_ids, 'attention_mask': attention_mask, 'output_hidden_states': True, 'return_dict': True}
52
+ if pixel_values is not None: bkw['pixel_values'] = pixel_values
53
+ if image_grid_thw is not None: bkw['image_grid_thw'] = image_grid_thw
54
+ out = self.backbone(**bkw)
55
+ h = out.hidden_states[-1]
56
+ if attention_mask is not None:
57
+ sl = attention_mask.sum(dim=1) - 1
58
+ feat = h[torch.arange(h.shape[0], device=h.device), sl]
59
+ else:
60
+ feat = h[:, -1, :]
61
+ pred = self.action_head(feat.float())
62
+ if action_targets is not None:
63
+ return F.l1_loss(pred, action_targets.float()), pred
64
+ return pred
65
+
66
+ def save_pretrained(self, save_dir):
67
+ os.makedirs(save_dir, exist_ok=True)
68
+ self.backbone.save_pretrained(save_dir)
69
+ torch.save(self.action_head.state_dict(), os.path.join(save_dir, "action_head.pt"))
70
+ with open(os.path.join(save_dir, "vla_config.json"), 'w') as f:
71
+ json.dump({'hidden_size': self.hidden_size, 'action_dim': 4, 'action_labels': ['Vx','Vy','Vz','Yaw_rate']}, f)
72
+
73
+ class VLADataset(TorchDataset):
74
+ def __init__(self, hf_dataset, processor, max_length=512):
75
+ self.dataset, self.processor, self.max_length = hf_dataset, processor, max_length
76
+ def __len__(self): return len(self.dataset)
77
+ def __getitem__(self, idx):
78
+ s = self.dataset[idx]
79
+ img = s['image'] if isinstance(s['image'], Image.Image) else Image.fromarray(np.array(s['image']))
80
+ img = img.convert('RGB')
81
+ action = torch.tensor(s['action'], dtype=torch.float32)
82
+ inst = s.get('language_instruction','navigate forward') or 'navigate forward'
83
+ msgs = [{"role":"user","content":[{"type":"image","image":img},{"type":"text","text":inst}]}]
84
+ try:
85
+ text = self.processor.apply_chat_template(msgs, tokenize=False, add_generation_prompt=True)
86
+ inputs = self.processor(text=[text], images=[img], padding='max_length', max_length=self.max_length, truncation=True, return_tensors="pt")
87
+ except:
88
+ inputs = self.processor(text=[inst], images=[img], padding='max_length', max_length=self.max_length, truncation=True, return_tensors="pt")
89
+ result = {k: v.squeeze(0) for k, v in inputs.items()}
90
+ result['action_targets'] = action
91
+ return result
92
+
93
+ def collate_fn(batch):
94
+ keys = batch[0].keys()
95
+ c = {}
96
+ for key in keys:
97
+ vals = [b[key] for b in batch]
98
+ if isinstance(vals[0], torch.Tensor):
99
+ if vals[0].dim() >= 1 and key in ['input_ids','attention_mask']:
100
+ ml = max(v.shape[0] for v in vals)
101
+ c[key] = torch.stack([F.pad(v,(0,ml-v.shape[0]),value=0) for v in vals])
102
+ else:
103
+ try: c[key] = torch.stack(vals)
104
+ except: c[key] = vals
105
+ else: c[key] = vals
106
+ return c
107
+
108
+ def train(args):
109
+ device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
110
+ print(f"Device: {device}")
111
+ ds = load_dataset(args.dataset_repo, split='train') if args.dataset_repo else load_from_disk(args.dataset_path)
112
+ print(f"Dataset: {len(ds)} rows")
113
+ processor = AutoProcessor.from_pretrained(args.model_name, trust_remote_code=True, min_pixels=224*224, max_pixels=512*512)
114
+ if processor.tokenizer.pad_token is None: processor.tokenizer.pad_token = processor.tokenizer.eos_token
115
+ model = DroneVLA(args.model_name, 4, args.lora_r, args.lora_alpha, args.freeze_vision).to(device)
116
+ loader = DataLoader(VLADataset(ds, processor, args.max_length), batch_size=args.batch_size, shuffle=True, num_workers=args.num_workers, collate_fn=collate_fn, pin_memory=device.type=='cuda', drop_last=True)
117
+ lora_p = [p for n,p in model.backbone.named_parameters() if p.requires_grad]
118
+ head_p = list(model.action_head.parameters())
119
+ opt = torch.optim.AdamW([{'params':lora_p,'lr':args.lr_lora,'weight_decay':0.01},{'params':head_p,'lr':args.lr_head}], betas=(0.9,0.95))
120
+ total = len(loader)*args.num_epochs; warmup = min(100, total//10)
121
+ sched = torch.optim.lr_scheduler.LambdaLR(opt, lambda s: s/max(warmup,1) if s<warmup else 0.1+0.9*(1+math.cos(math.pi*(s-warmup)/max(total-warmup,1)))/2)
122
+
123
+ if HAS_TRACKIO and args.hub_model_id:
124
+ try:
125
+ trackio.init(project="drone-vla", name=args.hub_model_id.split('/')[-1], space_id=f"{args.hub_model_id.split('/')[0]}/drone-vla-trackio")
126
+ except: pass
127
+
128
+ print(f"\nTraining: {args.num_epochs} epochs, batch={args.batch_size}, steps={total}")
129
+ step, best = 0, float('inf')
130
+ for epoch in range(args.num_epochs):
131
+ model.train(); el, es = 0, 0
132
+ for batch in loader:
133
+ ids = batch['input_ids'].to(device); mask = batch['attention_mask'].to(device); tgt = batch['action_targets'].to(device)
134
+ pv = batch.get('pixel_values')
135
+ if pv is not None: pv = (torch.stack(pv) if isinstance(pv,list) else pv).to(device, dtype=torch.bfloat16)
136
+ igt = batch.get('image_grid_thw')
137
+ if igt is not None: igt = (torch.stack(igt) if isinstance(igt,list) else igt).to(device)
138
+ loss, pred = model(ids, mask, pv, igt, tgt)
139
+ loss.backward(); torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
140
+ opt.step(); sched.step(); opt.zero_grad()
141
+ lv = loss.item(); el += lv; es += 1; step += 1
142
+ with torch.no_grad(): de = (pred-tgt).abs().mean(0)
143
+ if step % args.log_every == 0 or step == 1:
144
+ print(f"Step {step:5d} | Ep {epoch+1}/{args.num_epochs} | Loss: {lv:.4f} | Avg: {el/es:.4f} | Vx={de[0]:.4f} Vy={de[1]:.4f} Vz={de[2]:.4f} Yaw={de[3]:.4f}")
145
+ if HAS_TRACKIO:
146
+ try: trackio.log({"train/loss":lv,"train/avg_loss":el/es,"action_error/Vx":de[0].item(),"action_error/Vy":de[1].item(),"action_error/Vz":de[2].item(),"action_error/Yaw":de[3].item()})
147
+ except: pass
148
+ if step % args.save_every == 0:
149
+ model.save_pretrained(f"{args.output_dir}/ckpt-{step}")
150
+ if lv < best: best=lv; model.save_pretrained(f"{args.output_dir}/best"); print(f" ⭐ best: {best:.4f}")
151
+ print(f"\nEpoch {epoch+1} done | Avg: {el/max(es,1):.4f}\n")
152
+
153
+ final = f"{args.output_dir}/final"; model.save_pretrained(final)
154
+ if args.hub_model_id:
155
+ try:
156
+ model.backbone.push_to_hub(args.hub_model_id)
157
+ from huggingface_hub import HfApi; api = HfApi()
158
+ api.upload_file(path_or_fileobj=f"{final}/action_head.pt", path_in_repo="action_head.pt", repo_id=args.hub_model_id)
159
+ api.upload_file(path_or_fileobj=f"{final}/vla_config.json", path_in_repo="vla_config.json", repo_id=args.hub_model_id)
160
+ print(f"✅ https://huggingface.co/{args.hub_model_id}")
161
+ except Exception as e: print(f"Push failed: {e}")
162
+ if HAS_TRACKIO:
163
+ try: trackio.finish()
164
+ except: pass
165
+ print(f"Done! Best loss: {best:.4f}")
166
+
167
+ def main():
168
+ p = argparse.ArgumentParser()
169
+ p.add_argument('--dataset_repo', type=str, default=None)
170
+ p.add_argument('--dataset_path', type=str, default=None)
171
+ p.add_argument('--model_name', type=str, default='Qwen/Qwen3.5-0.8B')
172
+ p.add_argument('--lora_r', type=int, default=16)
173
+ p.add_argument('--lora_alpha', type=int, default=32)
174
+ p.add_argument('--freeze_vision', action='store_true', default=True)
175
+ p.add_argument('--max_length', type=int, default=512)
176
+ p.add_argument('--num_epochs', type=int, default=10)
177
+ p.add_argument('--batch_size', type=int, default=4)
178
+ p.add_argument('--lr_lora', type=float, default=2e-4)
179
+ p.add_argument('--lr_head', type=float, default=1e-3)
180
+ p.add_argument('--num_workers', type=int, default=2)
181
+ p.add_argument('--log_every', type=int, default=10)
182
+ p.add_argument('--save_every', type=int, default=500)
183
+ p.add_argument('--output_dir', type=str, default='./checkpoints')
184
+ p.add_argument('--hub_model_id', type=str, default=None)
185
+ a = p.parse_args()
186
+ if not a.dataset_repo and not a.dataset_path: print("Need --dataset_repo or --dataset_path"); sys.exit(1)
187
+ train(a)
188
+
189
+ if __name__ == '__main__': main()