Add training script
Browse files- 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()
|