Download utils/translation.py from plice13/SLT-space: direct link, hf CLI and curl.
- Browser
- Download file 1.21 kB
-
https://huggingface.co/spaces/plice13/SLT-space/resolve/c4e019a5171330dc35b2328bb7dce9e33650b93c/utils/translation.py
- Command line
-
hf download hf://spaces/plice13/SLT-space@c4e019a5171330dc35b2328bb7dce9e33650b93c/utils/translation.py
-
curl -L -o translation.py https://huggingface.co/spaces/plice13/SLT-space/resolve/c4e019a5171330dc35b2328bb7dce9e33650b93c/utils/translation.py
1.21 kB
| import torch | |
| def postprocess_text(preds, labels): | |
| preds = [pred.strip() for pred in preds] | |
| labels = [[label.strip()] for label in labels] | |
| return preds, labels | |
| # Add collate_fn to DataLoader | |
| def collate_fn(batch): | |
| # Add padding to the inputs | |
| # "inputs" must be 250 frames long | |
| # "attention_mask" must be 250 frames long | |
| # "labels" must be 128 tokens long | |
| return { | |
| "sign_inputs": torch.stack([ | |
| torch.cat((sample["sign_inputs"], torch.zeros(250 - sample["sign_inputs"].shape[0], 208)), dim=0) | |
| for sample in batch | |
| ]), | |
| "attention_mask": torch.stack([ | |
| torch.cat((sample["attention_mask"], torch.zeros(250 - sample["attention_mask"].shape[0])), dim=0) | |
| if sample["attention_mask"].shape[0] < 250 | |
| else sample["attention_mask"] | |
| for sample in batch | |
| ]), | |
| "labels": torch.stack([ | |
| torch.cat((sample["labels"].squeeze(0), torch.zeros(128 - sample["labels"].shape[0])), dim=0) | |
| if sample["labels"].shape[0] < 128 | |
| else sample["labels"] | |
| for sample in batch | |
| ]).squeeze(0).to(torch.long), | |
| } | |