dmusingu commited on
Commit
06f5b66
·
verified ·
1 Parent(s): a336f59

Update README with model loading code

Browse files
Files changed (1) hide show
  1. README.md +24 -6
README.md CHANGED
@@ -13,12 +13,30 @@ Part of the [LAPVQA collection](https://huggingface.co/collections/dmusingu/lapv
13
 
14
  ## Description
15
 
16
- RRG decoder trained on top of the **LAPVQA captioning-pretrained encoder**
17
  ([`lapvqa-pretrain-captioning`](https://huggingface.co/dmusingu/lapvqa-pretrain-captioning)).
18
- The encoder is kept frozen during RRG training; this file contains the decoder only.
19
 
20
- ## Files
21
 
22
- | File | Description |
23
- |---|---|
24
- | `pretrain-captioning.pt` | RRG decoder weights (encoder not included) |
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
13
 
14
  ## Description
15
 
16
+ RRG decoder trained on the frozen **LAPVQA captioning-pretrained encoder**
17
  ([`lapvqa-pretrain-captioning`](https://huggingface.co/dmusingu/lapvqa-pretrain-captioning)).
18
+ Checkpoint format: `{state_dict, vis_dim, d_model, num_layers, nhead, encoder, epoch, val_bleu4}`.
19
 
20
+ ## Loading
21
 
22
+ ```python
23
+ import torch
24
+ import tiktoken
25
+ from lapvqa.rrg.heads import ReportGenerationHead
26
+
27
+ ckpt = torch.load("pretrain-captioning.pt", map_location="cpu")
28
+ head = ReportGenerationHead(
29
+ vis_dim = ckpt["vis_dim"],
30
+ d_model = ckpt["d_model"],
31
+ num_layers = ckpt["num_layers"],
32
+ nhead = ckpt["nhead"],
33
+ )
34
+ head.load_state_dict(ckpt["state_dict"])
35
+ head.eval()
36
+
37
+ enc = tiktoken.get_encoding("gpt2")
38
+ bos_id = eos_id = enc.eot_token
39
+ # pair with encoder_final.pt from lapvqa-pretrain-captioning
40
+ token_ids = head.generate(vis_tokens, bos_id=bos_id, eos_id=eos_id)
41
+ reports = [enc.decode(ids) for ids in token_ids]
42
+ ```