update video caption part
This commit is contained in:
@@ -0,0 +1,174 @@
|
||||
# Video Caption
|
||||
English | [简体中文](./README_zh-CN.md)
|
||||
|
||||
The folder contains codes for dataset preprocessing (i.e., video splitting, filtering, and recaptioning), and beautiful prompt used by CogVideoX-Fun.
|
||||
The entire process supports distributed parallel processing, capable of handling large-scale datasets.
|
||||
|
||||
Meanwhile, we are collaborating with [Data-Juicer](https://github.com/modelscope/data-juicer/blob/main/docs/DJ_SORA.md),
|
||||
allowing you to easily perform video data processing on [Aliyun PAI-DLC](https://help.aliyun.com/zh/pai/user-guide/video-preprocessing/).
|
||||
|
||||
# Table of Content
|
||||
- [Video Caption](#video-caption)
|
||||
- [Table of Content](#table-of-content)
|
||||
- [Quick Start](#quick-start)
|
||||
- [Setup](#setup)
|
||||
- [Data Preprocessing](#data-preprocessing)
|
||||
- [Data Preparation](#data-preparation)
|
||||
- [Video Splitting](#video-splitting)
|
||||
- [Video Filtering](#video-filtering)
|
||||
- [Video Recaptioning](#video-recaptioning)
|
||||
- [Beautiful Prompt (For CogVideoX-Fun Inference)](#beautiful-prompt-for-cogvideox-inference)
|
||||
- [Batched Inference](#batched-inference)
|
||||
- [OpenAI Server](#openai-server)
|
||||
|
||||
## Quick Start
|
||||
|
||||
### Setup
|
||||
AliyunDSW or Docker is recommended to setup the environment, please refer to [Quick Start](../../README.md#quick-start).
|
||||
You can also refer to the image build process in the [Dockerfile](../../Dockerfile.ds) to configure the conda environment and other dependencies locally.
|
||||
|
||||
Since the video recaptioning depends on [llm-awq](https://github.com/mit-han-lab/llm-awq) for faster and memory efficient inference,
|
||||
the minimum GPU requirment should be RTX 3060 or A2 (CUDA Compute Capability >= 8.0).
|
||||
|
||||
```shell
|
||||
# pull image
|
||||
docker pull mybigpai-public-registry.cn-beijing.cr.aliyuncs.com/easycv/torch_cuda:cogvideox_fun
|
||||
|
||||
# enter image
|
||||
docker run -it -p 7860:7860 --network host --gpus all --security-opt seccomp:unconfined --shm-size 200g mybigpai-public-registry.cn-beijing.cr.aliyuncs.com/easycv/torch_cuda:cogvideox_fun
|
||||
|
||||
# clone code
|
||||
git clone https://github.com/aigc-apps/CogVideoX-Fun.git
|
||||
|
||||
# enter video_caption
|
||||
cd CogVideoX-Fun/cogvideox/video_caption
|
||||
```
|
||||
|
||||
### Data Preprocessing
|
||||
#### Data Preparation
|
||||
Place the downloaded videos into a folder under [datasets](./datasets/) (preferably without nested structures, as the video names are used as unique IDs in subsequent processes).
|
||||
Taking Panda-70M as an example, the entire dataset directory structure is shown as follows:
|
||||
```
|
||||
📦 datasets/
|
||||
├── 📂 panda_70m/
|
||||
│ ├── 📂 videos/
|
||||
│ │ ├── 📂 data/
|
||||
│ │ │ └── 📄 --C66yU3LjM_2.mp4
|
||||
│ │ │ └── 📄 ...
|
||||
```
|
||||
|
||||
#### Video Splitting
|
||||
CogVideoX-Fun utilizes [PySceneDetect](https://github.com/Breakthrough/PySceneDetect) to identify scene changes within the video
|
||||
and performs video splitting via FFmpeg based on certain threshold values to ensure consistency of the video clip.
|
||||
Video clips shorter than 3 seconds will be discarded, and those longer than 10 seconds will be splitted recursively.
|
||||
|
||||
The entire workflow of video splitting is in the [stage_1_video_splitting.sh](./scripts/stage_1_video_splitting.sh).
|
||||
After running
|
||||
```shell
|
||||
sh scripts/stage_1_video_splitting.sh
|
||||
```
|
||||
the video clips are obtained in `cogvideox/video_caption/datasets/panda_70m/videos_clips/data/`.
|
||||
|
||||
#### Video Filtering
|
||||
Based on the videos obtained in the previous step, CogVideoX-Fun provides a simple yet effective pipeline to filter out high-quality videos for recaptioning.
|
||||
The overall process is as follows:
|
||||
|
||||
- Aesthetic filtering: Filter out videos with poor content (blurry, dim, etc.) by calculating the average aesthetic score of uniformly sampled 4 frames via [aesthetic-predictor-v2-5](https://github.com/discus0434/aesthetic-predictor-v2-5).
|
||||
- Text filtering: Use [EasyOCR](https://github.com/JaidedAI/EasyOCR) to calculate the text area proportion of the middle frame to filter out videos with a large area of text.
|
||||
- Motion filtering: Calculate interframe optical flow differences to filter out videos that move too slowly or too quickly.
|
||||
|
||||
The entire workflow of video filtering is in the [stage_2_video_filtering.sh](./scripts/stage_2_video_filtering.sh).
|
||||
After running
|
||||
```shell
|
||||
sh scripts/stage_2_video_filtering.sh
|
||||
```
|
||||
the aesthetic score, text score, and motion score of videos will be saved in the corresponding meta files in the folder `cogvideox/video_caption/datasets/panda_70m/videos_clips/`.
|
||||
|
||||
> [!NOTE]
|
||||
> The computation of the aesthetic score depends on the [google/siglip-so400m-patch14-384 model](https://huggingface.co/google/siglip-so400m-patch14-384).
|
||||
Please run `HF_ENDPOINT=https://hf-mirror.com sh scripts/stage_2_video_filtering.sh` if you cannot access to huggingface.com.
|
||||
|
||||
|
||||
#### Video Recaptioning
|
||||
After obtaining the aboved high-quality filtered videos, CogVideoX-Fun utilizes [VILA1.5](https://github.com/NVlabs/VILA) to perform video recaptioning.
|
||||
Subsequently, the recaptioning results are rewritten by LLMs to better meet with the requirements of video generation tasks.
|
||||
Finally, an advanced VideoCLIPXL model is developed to filter out video-caption pairs with poor alignment, resulting in the final training dataset.
|
||||
|
||||
Please download the video caption model from [VILA1.5](https://huggingface.co/collections/Efficient-Large-Model/vila-on-pre-training-for-visual-language-models-65d8022a3a52cd9bcd62698e) of the appropriate size based on the GPU memory of your machine.
|
||||
For A100 with 40G VRAM, you can download [VILA1.5-40b-AWQ](https://huggingface.co/Efficient-Large-Model/VILA1.5-40b-AWQ) by running
|
||||
```shell
|
||||
# Add HF_ENDPOINT=https://hf-mirror.com before the command if you cannot access to huggingface.com
|
||||
huggingface-cli download Efficient-Large-Model/VILA1.5-40b-AWQ --local-dir-use-symlinks False --local-dir /PATH/TO/VILA_MODEL
|
||||
```
|
||||
|
||||
Optionally, you can prepare local LLMs to rewrite the recaption results.
|
||||
For example, you can download [Meta-Llama-3-8B-Instruct](https://huggingface.co/NousResearch/Meta-Llama-3-8B-Instruct) by running
|
||||
```shell
|
||||
# Add HF_ENDPOINT=https://hf-mirror.com before the command if you cannot access to huggingface.com
|
||||
huggingface-cli download NousResearch/Meta-Llama-3-8B-Instruct --local-dir-use-symlinks False --local-dir /PATH/TO/REWRITE_MODEL
|
||||
```
|
||||
|
||||
The entire workflow of video recaption is in the [stage_3_video_recaptioning.sh](./scripts/stage_3_video_recaptioning.sh).
|
||||
After running
|
||||
```shell
|
||||
VILA_MODEL_PATH=/PATH/TO/VILA_MODEL REWRITE_MODEL_PATH=/PATH/TO/REWRITE_MODEL sh scripts/stage_3_video_recaptioning.sh
|
||||
```
|
||||
the final train file is obtained in `cogvideox/video_caption/datasets/panda_70m/videos_clips/meta_train_info.json`.
|
||||
|
||||
|
||||
### Beautiful Prompt (For CogVideoX-Fun Inference)
|
||||
Beautiful Prompt aims to rewrite and beautify the user-uploaded prompt via LLMs, mapping it to the style of CogVideoX-Fun's training captions,
|
||||
making it more suitable as the inference prompt and thus improving the quality of the generated videos.
|
||||
We support batched inference with local LLMs or OpenAI compatible server based on [vLLM](https://github.com/vllm-project/vllm) for beautiful prompt.
|
||||
|
||||
#### Batched Inference
|
||||
1. Prepare original prompts in a jsonl file `cogvideox/video_caption/datasets/original_prompt.jsonl` with the following format:
|
||||
```json
|
||||
{"prompt": "A stylish woman in a black leather jacket, red dress, and boots walks confidently down a damp Tokyo street."}
|
||||
{"prompt": "An underwater world with realistic fish and other creatures of the sea."}
|
||||
{"prompt": "a monarch butterfly perched on a tree trunk in the forest."}
|
||||
{"prompt": "a child in a room with a bottle of wine and a lamp."}
|
||||
{"prompt": "two men in suits walking down a hallway."}
|
||||
```
|
||||
|
||||
2. Then you can perform beautiful prompt by running
|
||||
```shell
|
||||
# Meta-Llama-3-8B-Instruct is sufficient for this task.
|
||||
# Download it from https://huggingface.co/NousResearch/Meta-Llama-3-8B-Instruct or https://www.modelscope.cn/models/LLM-Research/Meta-Llama-3-8B-Instruct to /path/to/your_llm
|
||||
|
||||
python caption_rewrite.py \
|
||||
--video_metadata_path datasets/original_prompt.jsonl \
|
||||
--caption_column "prompt" \
|
||||
--batch_size 1 \
|
||||
--model_name /path/to/your_llm \
|
||||
--prompt prompt/beautiful_prompt.txt \
|
||||
--prefix '"detailed description": ' \
|
||||
--saved_path datasets/beautiful_prompt.jsonl \
|
||||
--saved_freq 1
|
||||
```
|
||||
|
||||
#### OpenAI Server
|
||||
+ You can request OpenAI compatible server to perform beautiful prompt by running
|
||||
```shell
|
||||
OPENAI_API_KEY="your_openai_api_key" OPENAI_BASE_URL="your_openai_base_url" python beautiful_prompt.py \
|
||||
--model "your_model_name" \
|
||||
--prompt "your_prompt"
|
||||
```
|
||||
|
||||
+ You can also deploy the OpenAI Compatible Server locally using vLLM. For example:
|
||||
```shell
|
||||
# Meta-Llama-3-8B-Instruct is sufficient for this task.
|
||||
# Download it from https://huggingface.co/NousResearch/Meta-Llama-3-8B-Instruct or https://www.modelscope.cn/models/LLM-Research/Meta-Llama-3-8B-Instruct to /path/to/your_llm
|
||||
|
||||
# deploy the OpenAI compatible server
|
||||
python -m vllm.entrypoints.openai.api_server serve /path/to/your_llm --dtype auto --api-key "your_api_key"
|
||||
```
|
||||
|
||||
Then you can perform beautiful prompt by running
|
||||
```shell
|
||||
python -m beautiful_prompt.py \
|
||||
--model /path/to/your_llm \
|
||||
--prompt "your_prompt" \
|
||||
--base_url "http://localhost:8000/v1" \
|
||||
--api_key "your_api_key"
|
||||
```
|
||||
@@ -0,0 +1,159 @@
|
||||
# 数据预处理
|
||||
[English](./README.md) | 简体中文
|
||||
|
||||
该文件夹包含 CogVideoX-Fun 使用的数据集预处理(即视频切分、过滤和生成描述)和提示词美化的代码。整个过程支持分布式并行处理,能够处理大规模数据集。
|
||||
|
||||
此外,我们和 [Data-Juicer](https://github.com/modelscope/data-juicer/blob/main/docs/DJ_SORA.md) 合作,能让你在 [Aliyun PAI-DLC](https://help.aliyun.com/zh/pai/user-guide/video-preprocessing/) 轻松进行视频数据的处理。
|
||||
|
||||
# 目录
|
||||
- [数据预处理](#数据预处理)
|
||||
- [目录](#目录)
|
||||
- [快速开始](#快速开始)
|
||||
- [安装](#安装)
|
||||
- [数据集预处理](#数据集预处理)
|
||||
- [数据准备](#数据准备)
|
||||
- [视频切分](#视频切分)
|
||||
- [视频过滤](#视频过滤)
|
||||
- [视频描述](#视频描述)
|
||||
- [提示词美化](#提示词美化)
|
||||
- [批量推理](#批量推理)
|
||||
- [OpenAI 服务器](#openai-服务器)
|
||||
|
||||
|
||||
## 快速开始
|
||||
### 安装
|
||||
推荐使用阿里云 DSW 和 Docker 来安装环境,请参考 [快速开始](../../README_zh-CN.md#1-云使用-aliyundswdocker). 你也可以参考 [Dockerfile](../../Dockerfile.ds) 中的镜像构建流程在本地安装对应的 conda 环境和其余依赖。
|
||||
|
||||
为了提高推理速度和节省推理的显存,生成视频描述依赖于 [llm-awq](https://github.com/mit-han-lab/llm-awq)。因此,需要 RTX 3060 或者 A2 及以上的显卡 (CUDA Compute Capability >= 8.0)。
|
||||
|
||||
```shell
|
||||
# pull image
|
||||
docker pull mybigpai-public-registry.cn-beijing.cr.aliyuncs.com/easycv/torch_cuda:cogvideox_fun
|
||||
|
||||
# enter image
|
||||
docker run -it -p 7860:7860 --network host --gpus all --security-opt seccomp:unconfined --shm-size 200g mybigpai-public-registry.cn-beijing.cr.aliyuncs.com/easycv/torch_cuda:cogvideox_fun
|
||||
|
||||
# clone code
|
||||
git clone https://github.com/aigc-apps/CogVideoX-Fun.git
|
||||
|
||||
# enter video_caption
|
||||
cd CogVideoX-Fun/cogvideox/video_caption
|
||||
```
|
||||
|
||||
### 数据集预处理
|
||||
#### 数据准备
|
||||
将下载的视频准备到文件夹 [datasets](./datasets/)(最好不使用嵌套结构,因为视频名称在后续处理中用作唯一 ID)。以 Panda-70M 为例,完整的数据集目录结构如下所示:
|
||||
```
|
||||
📦 datasets/
|
||||
├── 📂 panda_70m/
|
||||
│ ├── 📂 videos/
|
||||
│ │ ├── 📂 data/
|
||||
│ │ │ └── 📄 --C66yU3LjM_2.mp4
|
||||
│ │ │ └── 📄 ...
|
||||
```
|
||||
|
||||
#### 视频切分
|
||||
CogVideoX-Fun 使用 [PySceneDetect](https://github.com/Breakthrough/PySceneDetect) 来识别视频中的场景变化
|
||||
并根据某些阈值通过 FFmpeg 执行视频分割,以确保视频片段的一致性。
|
||||
短于 3 秒的视频片段将被丢弃,长于 10 秒的视频片段将被递归切分。
|
||||
|
||||
视频切分的完整流程在 [stage_1_video_splitting.sh](./scripts/stage_1_video_splitting.sh)。执行
|
||||
```shell
|
||||
sh scripts/stage_1_video_splitting.sh
|
||||
```
|
||||
后,切分后的视频位于 `cogvideox/video_caption/datasets/panda_70m/videos_clips/data/`。
|
||||
|
||||
#### 视频过滤
|
||||
基于上一步获得的视频,CogVideoX-Fun 提供了一个简单而有效的流程来过滤出高质量的视频。总体流程如下:
|
||||
|
||||
- 美学过滤:通过 [aesthetic-predictor-v2-5](https://github.com/discus0434/aesthetic-predictor-v2-5) 计算均匀采样的 4 帧视频的平均美学分数,从而筛选出内容不佳(模糊、昏暗等)的视频。
|
||||
- 文本过滤:使用 [EasyOCR](https://github.com/JaidedAI/EasyOCR) 计算中间帧的文本区域比例,过滤掉含有大面积文本的视频。
|
||||
- 运动过滤:计算帧间光流差,过滤掉移动太慢或太快的视频。
|
||||
|
||||
视频过滤的完整流程在 [stage_2_video_filtering.sh](./scripts/stage_2_video_filtering.sh)。执行
|
||||
```shell
|
||||
sh scripts/stage_2_video_filtering.sh
|
||||
```
|
||||
后,视频的美学得分、文本得分和运动得分对应的元文件保存在 `cogvideox/video_caption/datasets/panda_70m/videos_clips/`。
|
||||
|
||||
> [!NOTE]
|
||||
> 美学得分的计算依赖于 [google/siglip-so400m-patch14-384 model](https://huggingface.co/google/siglip-so400m-patch14-384).
|
||||
请执行 `HF_ENDPOINT=https://hf-mirror.com sh scripts/stage_2_video_filtering.sh` 如果你无法访问 huggingface.com.
|
||||
|
||||
#### 视频描述
|
||||
在获得上述高质量的过滤视频后,CogVideoX-Fun 利用 [VILA1.5](https://github.com/NVlabs/VILA) 来生成视频描述。随后,使用 LLMs 对生成的视频描述进行重写,以更好地满足视频生成任务的要求。最后,使用自研的 VideoCLIPXL 模型来过滤掉描述和视频内容不一致的数据,从而得到最终的训练数据集。
|
||||
|
||||
请根据机器的显存从 [VILA1.5](https://huggingface.co/collections/Efficient-Large-Model/vila-on-pre-training-for-visual-language-models-65d8022a3a52cd9bcd62698e) 下载合适大小的模型。对于 A100 40G,你可以执行下面的命令来下载 [VILA1.5-40b-AWQ](https://huggingface.co/Efficient-Large-Model/VILA1.5-40b-AWQ)
|
||||
```shell
|
||||
# Add HF_ENDPOINT=https://hf-mirror.com before the command if you cannot access to huggingface.com
|
||||
huggingface-cli download Efficient-Large-Model/VILA1.5-40b-AWQ --local-dir-use-symlinks False --local-dir /PATH/TO/VILA_MODEL
|
||||
```
|
||||
|
||||
你可以选择性地准备 LLMs 来改写上述视频描述的结果。例如,你执行下面的命令来下载 [Meta-Llama-3-8B-Instruct](https://huggingface.co/NousResearch/Meta-Llama-3-8B-Instruct)
|
||||
```shell
|
||||
# Add HF_ENDPOINT=https://hf-mirror.com before the command if you cannot access to huggingface.com
|
||||
huggingface-cli download NousResearch/Meta-Llama-3-8B-Instruct --local-dir-use-symlinks False --local-dir /PATH/TO/REWRITE_MODEL
|
||||
```
|
||||
|
||||
视频描述的完整流程在 [stage_3_video_recaptioning.sh](./scripts/stage_3_video_recaptioning.sh).
|
||||
执行
|
||||
```shell
|
||||
VILA_MODEL_PATH=/PATH/TO/VILA_MODEL REWRITE_MODEL_PATH=/PATH/TO/REWRITE_MODEL sh scripts/stage_3_video_recaptioning.sh
|
||||
```
|
||||
后,最后的训练文件会保存在 `cogvideox/video_caption/datasets/panda_70m/videos_clips/meta_train_info.json`。
|
||||
|
||||
### 提示词美化
|
||||
提示词美化旨在通过 LLMs 重写和美化用户上传的提示,将其映射为 CogVideoX-Fun 训练所使用的视频描述风格、
|
||||
使其更适合用作推理提示词,从而提高生成视频的质量。
|
||||
|
||||
基于 [vLLM](https://github.com/vllm-project/vllm),我们支持使用本地 LLM 进行批量推理或请求 OpenAI 服务器的方式,以进行提示词美化。
|
||||
|
||||
#### 批量推理
|
||||
1. 将原始的提示词以下面的格式准备在文件 `cogvideox/video_caption/datasets/original_prompt.jsonl` 中:
|
||||
```json
|
||||
{"prompt": "A stylish woman in a black leather jacket, red dress, and boots walks confidently down a damp Tokyo street."}
|
||||
{"prompt": "An underwater world with realistic fish and other creatures of the sea."}
|
||||
{"prompt": "a monarch butterfly perched on a tree trunk in the forest."}
|
||||
{"prompt": "a child in a room with a bottle of wine and a lamp."}
|
||||
{"prompt": "two men in suits walking down a hallway."}
|
||||
```
|
||||
|
||||
2. 随后你可以通过执行以下的命令进行提示词美化
|
||||
```shell
|
||||
# Meta-Llama-3-8B-Instruct is sufficient for this task.
|
||||
# Download it from https://huggingface.co/NousResearch/Meta-Llama-3-8B-Instruct or https://www.modelscope.cn/models/LLM-Research/Meta-Llama-3-8B-Instruct to /path/to/your_llm
|
||||
|
||||
python caption_rewrite.py \
|
||||
--video_metadata_path datasets/original_prompt.jsonl \
|
||||
--caption_column "prompt" \
|
||||
--batch_size 1 \
|
||||
--model_name /path/to/your_llm \
|
||||
--prompt prompt/beautiful_prompt.txt \
|
||||
--prefix '"detailed description": ' \
|
||||
--saved_path datasets/beautiful_prompt.jsonl \
|
||||
--saved_freq 1
|
||||
```
|
||||
|
||||
#### OpenAI 服务器
|
||||
+ 你可以通过请求 OpenAI 服务器的方式来进行提示词美化
|
||||
```shell
|
||||
OPENAI_API_KEY="your_openai_api_key" OPENAI_BASE_URL="your_openai_base_url" python beautiful_prompt.py \
|
||||
--model "your_model_name" \
|
||||
--prompt "your_prompt"
|
||||
```
|
||||
|
||||
+ 你也可以执行以下命令,通过 vLLM 将本地 LLMs 部署成兼容 OpenAI 的服务器
|
||||
```shell
|
||||
OPENAI_API_KEY="your_openai_api_key" OPENAI_BASE_URL="your_openai_base_url" python beautiful_prompt.py \
|
||||
--model "your_model_name" \
|
||||
--prompt "your_prompt"
|
||||
```
|
||||
|
||||
然后再执行下面的命令来进行提示词美化
|
||||
```shell
|
||||
python -m beautiful_prompt.py \
|
||||
--model /path/to/your_llm \
|
||||
--prompt "your_prompt" \
|
||||
--base_url "http://localhost:8000/v1" \
|
||||
--api_key "your_api_key"
|
||||
```
|
||||
@@ -0,0 +1,103 @@
|
||||
"""
|
||||
This script (optional) can rewrite and beautify the user-uploaded prompt via LLMs, mapping it to the style of cogvideox's training captions,
|
||||
making it more suitable as the inference prompt and thus improving the quality of the generated videos.
|
||||
|
||||
Usage:
|
||||
+ You can request OpenAI compatible server to perform beautiful prompt by running
|
||||
```shell
|
||||
export OPENAI_API_KEY="your_openai_api_key" OPENAI_BASE_URL="your_openai_base_url" python beautiful_prompt.py \
|
||||
--model "your_model_name" \
|
||||
--prompt "your_prompt"
|
||||
```
|
||||
+ You can also deploy the OpenAI Compatible Server locally using vLLM. For example:
|
||||
```shell
|
||||
# Meta-Llama-3-8B-Instruct is sufficient for this task.
|
||||
# Download it from https://huggingface.co/NousResearch/Meta-Llama-3-8B-Instruct or https://www.modelscope.cn/models/LLM-Research/Meta-Llama-3-8B-Instruct to /path/to/your_llm
|
||||
|
||||
# deploy the OpenAI compatible server
|
||||
python -m vllm.entrypoints.openai.api_server serve /path/to/your_llm --dtype auto --api-key "your_api_key"
|
||||
```
|
||||
|
||||
Then you can perform beautiful prompt by running
|
||||
```shell
|
||||
python -m beautiful_prompt.py \
|
||||
--model /path/to/your_llm \
|
||||
--prompt "your_prompt" \
|
||||
--base_url "http://localhost:8000/v1" \
|
||||
--api_key "your_api_key"
|
||||
```
|
||||
"""
|
||||
import argparse
|
||||
import os
|
||||
|
||||
from openai import OpenAI
|
||||
|
||||
from cogvideox.video_caption.caption_rewrite import extract_output
|
||||
|
||||
|
||||
def parse_args():
|
||||
parser = argparse.ArgumentParser(description="Beautiful prompt.")
|
||||
parser.add_argument("--model", type=str, required=True, help="The OpenAI model or the path to your local LLM.")
|
||||
parser.add_argument("--prompt", type=str, required=True, help="The user-uploaded prompt.")
|
||||
parser.add_argument(
|
||||
"--template",
|
||||
type=str,
|
||||
default="cogvideox/video_caption/prompt/beautiful_prompt.txt",
|
||||
help="A string or a txt file contains the template for beautiful prompt."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--max_retry_nums",
|
||||
type=int,
|
||||
default=5,
|
||||
help="Maximum number of retries to obtain an output that meets the JSON format."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--base_url",
|
||||
type=str,
|
||||
default=None,
|
||||
help="OpenAI API server url. If it is None, the OPENAI_BASE_URL from the environment variables will be used.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--api_key",
|
||||
type=str,
|
||||
default=None,
|
||||
help="OpenAI API key. If it is None, the OPENAI_API_KEY from the environment variables will be used.",
|
||||
)
|
||||
|
||||
args = parser.parse_args()
|
||||
return args
|
||||
|
||||
|
||||
def main():
|
||||
args = parse_args()
|
||||
|
||||
client = OpenAI(
|
||||
base_url=os.getenv("OPENAI_BASE_URL", args.base_url),
|
||||
api_key=os.environ.get("OPENAI_API_KEY", args.api_key),
|
||||
)
|
||||
if args.template.endswith(".txt") and os.path.exists(args.template):
|
||||
with open(args.template, "r") as f:
|
||||
args.template = "".join(f.readlines())
|
||||
# print(f"Beautiful prompt template: {args.template}")
|
||||
|
||||
for _ in range(args.max_retry_nums):
|
||||
completion = client.chat.completions.create(
|
||||
model=args.model,
|
||||
messages=[
|
||||
# {"role": "system", "content": "You are a helpful assistant."},
|
||||
{"role": "user", "content": args.template + "\n" + str(args.prompt)}
|
||||
],
|
||||
temperature=0.7,
|
||||
top_p=1,
|
||||
max_tokens=1024,
|
||||
)
|
||||
|
||||
output = completion.choices[0].message.content
|
||||
output = extract_output(output, prefix='"detailed description": ')
|
||||
if output is not None:
|
||||
break
|
||||
print(f"Beautiful prompt: {output}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,224 @@
|
||||
import argparse
|
||||
import re
|
||||
import os
|
||||
from tqdm import tqdm
|
||||
|
||||
import pandas as pd
|
||||
import torch
|
||||
from natsort import index_natsorted
|
||||
from vllm import LLM, SamplingParams
|
||||
from transformers import AutoTokenizer
|
||||
|
||||
from utils.logger import logger
|
||||
|
||||
|
||||
def extract_output(s, prefix='"rewritten description": '):
|
||||
"""Customize the function according to the prompt."""
|
||||
# Since some LLMs struggles to output strictly formatted JSON strings as specified by the prompt,
|
||||
# thus manually parse the output string `{"rewritten description": "your rewritten description here"}`.
|
||||
match = re.search(r"{(.+?)}", s, re.DOTALL)
|
||||
if not match:
|
||||
logger.warning(f"{s} is not in the json format. Return None.")
|
||||
return None
|
||||
output = match.group(1).strip()
|
||||
if output.startswith(prefix):
|
||||
output = output[len(prefix) :]
|
||||
if output[0] == '"' and output[-1] == '"':
|
||||
return output[1:-1]
|
||||
else:
|
||||
logger.warning(f"{output} does not start and end with the double quote. Return None.")
|
||||
return None
|
||||
else:
|
||||
logger.warning(f"{output} does not start with {prefix}. Return None.")
|
||||
return None
|
||||
|
||||
|
||||
def parse_args():
|
||||
parser = argparse.ArgumentParser(description="Rewrite the video caption by LLMs.")
|
||||
parser.add_argument(
|
||||
"--video_metadata_path", type=str, required=True, help="The path to the video dataset metadata (csv/jsonl)."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--video_path_column",
|
||||
type=str,
|
||||
default=None,
|
||||
help="The column contains the video path (an absolute path or a relative path w.r.t the video_folder).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--caption_column",
|
||||
type=str,
|
||||
default="caption",
|
||||
help="The column contains the video caption.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--batch_size",
|
||||
type=int,
|
||||
default=128,
|
||||
required=False,
|
||||
help="The batch size for vllm inference. Adjust according to the number of GPUs to maximize inference throughput.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--model_name",
|
||||
type=str,
|
||||
default="NousResearch/Meta-Llama-3-8B-Instruct",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--prompt",
|
||||
type=str,
|
||||
required=True,
|
||||
help="A string or a txt file contains the prompt.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--prefix",
|
||||
type=str,
|
||||
required=True,
|
||||
help="The prefix to extract the output from LLMs.",
|
||||
)
|
||||
parser.add_argument("--saved_path", type=str, required=True, help="The save path to the output results (csv/jsonl).")
|
||||
parser.add_argument("--saved_freq", type=int, default=1, help="The frequency to save the output results.")
|
||||
|
||||
args = parser.parse_args()
|
||||
return args
|
||||
|
||||
|
||||
def main():
|
||||
args = parse_args()
|
||||
|
||||
if args.video_metadata_path.endswith(".csv"):
|
||||
video_metadata_df = pd.read_csv(args.video_metadata_path)
|
||||
elif args.video_metadata_path.endswith(".jsonl"):
|
||||
video_metadata_df = pd.read_json(args.video_metadata_path, lines=True)
|
||||
elif args.video_metadata_path.endswith(".json"):
|
||||
video_metadata_df = pd.read_json(args.video_metadata_path)
|
||||
else:
|
||||
raise ValueError(f"The {args.video_metadata_path} must end with .csv, .jsonl or .json.")
|
||||
|
||||
saved_suffix = os.path.splitext(args.saved_path)[1]
|
||||
if saved_suffix not in set([".csv", ".jsonl", ".json"]):
|
||||
raise ValueError(f"The saved_path must end with .csv, .jsonl or .json.")
|
||||
|
||||
if os.path.exists(args.saved_path) and args.video_path_column is not None:
|
||||
if args.saved_path.endswith(".csv"):
|
||||
saved_metadata_df = pd.read_csv(args.saved_path)
|
||||
elif args.saved_path.endswith(".jsonl"):
|
||||
saved_metadata_df = pd.read_json(args.saved_path, lines=True)
|
||||
|
||||
# Filter out the unprocessed video-caption pairs by setting the indicator=True.
|
||||
merged_df = video_metadata_df.merge(saved_metadata_df, on=args.video_path_column, how="outer", indicator=True)
|
||||
video_metadata_df = merged_df[merged_df["_merge"] == "left_only"]
|
||||
# Sorting to guarantee the same result for each process.
|
||||
video_metadata_df = video_metadata_df.iloc[index_natsorted(video_metadata_df[args.video_path_column])].reset_index(
|
||||
drop=True
|
||||
)
|
||||
logger.info(
|
||||
f"Resume from {args.saved_path}: {len(saved_metadata_df)} processed and {len(video_metadata_df)} to be processed."
|
||||
)
|
||||
|
||||
if args.prompt.endswith(".txt") and os.path.exists(args.prompt):
|
||||
with open(args.prompt, "r") as f:
|
||||
args.prompt = "".join(f.readlines())
|
||||
logger.info(f"Prompt: {args.prompt}")
|
||||
|
||||
if args.video_path_column is not None:
|
||||
video_path_list = video_metadata_df[args.video_path_column].tolist()
|
||||
if args.caption_column in video_metadata_df.columns:
|
||||
sampled_frame_caption_list = video_metadata_df[args.caption_column].tolist()
|
||||
else:
|
||||
# When two columns with the same name, the dataframe merge operation on will distinguish them by adding 'x' and 'y'.
|
||||
sampled_frame_caption_list = video_metadata_df[args.caption_column + "_x"].tolist()
|
||||
|
||||
CUDA_VISIBLE_DEVICES = os.getenv("CUDA_VISIBLE_DEVICES", None)
|
||||
tensor_parallel_size = torch.cuda.device_count() if CUDA_VISIBLE_DEVICES is None else len(CUDA_VISIBLE_DEVICES.split(","))
|
||||
logger.info(f"Automatically set tensor_parallel_size={tensor_parallel_size} based on the available devices.")
|
||||
|
||||
llm = LLM(model=args.model_name, trust_remote_code=True, tensor_parallel_size=tensor_parallel_size)
|
||||
if "Meta-Llama-3" in args.model_name:
|
||||
if "Meta-Llama-3-70B" in args.model_name:
|
||||
# Llama-3-70B should use the tokenizer from Llama-3-8B
|
||||
# https://github.com/vllm-project/vllm/issues/4180#issuecomment-2068292942
|
||||
tokenizer = AutoTokenizer.from_pretrained("NousResearch/Meta-Llama-3-8B-Instruct")
|
||||
else:
|
||||
tokenizer = AutoTokenizer.from_pretrained(args.model_name)
|
||||
stop_token_ids = [tokenizer.eos_token_id, tokenizer.convert_tokens_to_ids("<|eot_id|>")]
|
||||
sampling_params = SamplingParams(temperature=0.7, top_p=1, max_tokens=1024, stop_token_ids=stop_token_ids)
|
||||
else:
|
||||
tokenizer = AutoTokenizer.from_pretrained(args.model_name)
|
||||
sampling_params = SamplingParams(temperature=0.7, top_p=1, max_tokens=1024)
|
||||
|
||||
result_dict = {args.caption_column: []}
|
||||
if args.video_path_column is not None:
|
||||
result_dict = {args.video_path_column: [], args.caption_column: []}
|
||||
|
||||
for i in tqdm(range(0, len(sampled_frame_caption_list), args.batch_size)):
|
||||
if args.video_path_column is not None:
|
||||
batch_video_path = video_path_list[i : i + args.batch_size]
|
||||
batch_caption = sampled_frame_caption_list[i : i + args.batch_size]
|
||||
batch_prompt = []
|
||||
for caption in batch_caption:
|
||||
# batch_prompt.append("user:" + args.prompt + str(caption) + "\n assistant:")
|
||||
messages = [
|
||||
{"role": "system", "content": "You are a helpful assistant."},
|
||||
{"role": "user", "content": args.prompt + "\n" + str(caption)},
|
||||
]
|
||||
text = tokenizer.apply_chat_template(messages, tokenize=False, add_generation_prompt=True)
|
||||
batch_prompt.append(text)
|
||||
|
||||
batch_output = llm.generate(batch_prompt, sampling_params)
|
||||
batch_output = [output.outputs[0].text.rstrip() for output in batch_output]
|
||||
batch_output = [extract_output(output, prefix=args.prefix) for output in batch_output]
|
||||
|
||||
# Filter out data that does not meet the output format.
|
||||
batch_result = []
|
||||
if args.video_path_column is not None:
|
||||
for video_path, output in zip(batch_video_path, batch_output):
|
||||
if output is not None:
|
||||
batch_result.append((video_path, output))
|
||||
batch_video_path, batch_output = zip(*batch_result)
|
||||
|
||||
result_dict[args.video_path_column].extend(batch_video_path)
|
||||
else:
|
||||
for output in batch_output:
|
||||
if output is not None:
|
||||
batch_result.append(output)
|
||||
|
||||
result_dict[args.caption_column].extend(batch_result)
|
||||
|
||||
# Save the metadata every args.saved_freq.
|
||||
if i != 0 and ((i // args.batch_size) % args.saved_freq) == 0:
|
||||
if len(result_dict[args.caption_column]) > 0:
|
||||
result_df = pd.DataFrame(result_dict)
|
||||
if args.saved_path.endswith(".csv"):
|
||||
header = True if not os.path.exists(args.saved_path) else False
|
||||
result_df.to_csv(args.saved_path, header=header, index=False, mode="a")
|
||||
elif args.saved_path.endswith(".jsonl"):
|
||||
result_df.to_json(args.saved_path, orient="records", lines=True, mode="a", force_ascii=False)
|
||||
elif args.saved_path.endswith(".json"):
|
||||
# Append is not supported.
|
||||
if os.path.exists(args.saved_path):
|
||||
saved_df = pd.read_json(args.saved_path, orient="records")
|
||||
result_df = pd.concat([saved_df, result_df], ignore_index=True)
|
||||
result_df.to_json(args.saved_path, orient="records", indent=4, force_ascii=False)
|
||||
logger.info(f"Save result to {args.saved_path}.")
|
||||
|
||||
result_dict = {args.caption_column: []}
|
||||
if args.video_path_column is not None:
|
||||
result_dict = {args.video_path_column: [], args.caption_column: []}
|
||||
|
||||
if len(result_dict[args.caption_column]) > 0:
|
||||
result_df = pd.DataFrame(result_dict)
|
||||
if args.saved_path.endswith(".csv"):
|
||||
header = True if not os.path.exists(args.saved_path) else False
|
||||
result_df.to_csv(args.saved_path, header=header, index=False, mode="a")
|
||||
elif args.saved_path.endswith(".jsonl"):
|
||||
result_df.to_json(args.saved_path, orient="records", lines=True, mode="a")
|
||||
elif args.saved_path.endswith(".json"):
|
||||
# Append is not supported.
|
||||
if os.path.exists(args.saved_path):
|
||||
saved_df = pd.read_json(args.saved_path, orient="records")
|
||||
result_df = pd.concat([saved_df, result_df], ignore_index=True)
|
||||
result_df.to_json(args.saved_path, orient="records", indent=4, force_ascii=False)
|
||||
logger.info(f"Save the final result to {args.saved_path}.")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,186 @@
|
||||
import ast
|
||||
import argparse
|
||||
import gc
|
||||
import os
|
||||
from contextlib import contextmanager
|
||||
from pathlib import Path
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
from joblib import Parallel, delayed
|
||||
from natsort import natsorted
|
||||
from tqdm import tqdm
|
||||
|
||||
from utils.logger import logger
|
||||
from utils.filter import filter
|
||||
|
||||
|
||||
@contextmanager
|
||||
def VideoCapture(video_path):
|
||||
cap = cv2.VideoCapture(video_path)
|
||||
try:
|
||||
yield cap
|
||||
finally:
|
||||
cap.release()
|
||||
del cap
|
||||
gc.collect()
|
||||
|
||||
|
||||
def compute_motion_score(video_path):
|
||||
video_motion_scores = []
|
||||
sampling_fps = 2
|
||||
|
||||
try:
|
||||
with VideoCapture(video_path) as cap:
|
||||
fps = cap.get(cv2.CAP_PROP_FPS)
|
||||
valid_fps = min(max(sampling_fps, 1), fps)
|
||||
frame_interval = int(fps / valid_fps)
|
||||
total_frames = int(cap.get(cv2.CAP_PROP_FRAME_COUNT))
|
||||
|
||||
# if cannot get the second frame, use the last one
|
||||
frame_interval = min(frame_interval, total_frames - 1)
|
||||
|
||||
prev_frame = None
|
||||
frame_count = -1
|
||||
while cap.isOpened():
|
||||
ret, frame = cap.read()
|
||||
frame_count += 1
|
||||
|
||||
if not ret:
|
||||
break
|
||||
|
||||
# skip middle frames
|
||||
if frame_count % frame_interval != 0:
|
||||
continue
|
||||
|
||||
gray_frame = cv2.cvtColor(frame, cv2.COLOR_BGR2GRAY)
|
||||
if prev_frame is None:
|
||||
prev_frame = gray_frame
|
||||
continue
|
||||
|
||||
flow = cv2.calcOpticalFlowFarneback(
|
||||
prev_frame,
|
||||
gray_frame,
|
||||
None,
|
||||
pyr_scale=0.5,
|
||||
levels=3,
|
||||
winsize=15,
|
||||
iterations=3,
|
||||
poly_n=5,
|
||||
poly_sigma=1.2,
|
||||
flags=0,
|
||||
)
|
||||
mag, _ = cv2.cartToPolar(flow[..., 0], flow[..., 1])
|
||||
frame_motion_score = np.mean(mag)
|
||||
video_motion_scores.append(frame_motion_score)
|
||||
prev_frame = gray_frame
|
||||
|
||||
video_meta_info = {
|
||||
"video_path": Path(video_path).name,
|
||||
"motion_score": round(float(np.mean(video_motion_scores)), 5),
|
||||
}
|
||||
return video_meta_info
|
||||
|
||||
except Exception as e:
|
||||
print(f"Compute motion score for video {video_path} with error: {e}.")
|
||||
|
||||
|
||||
def parse_args():
|
||||
parser = argparse.ArgumentParser(description="Compute the motion score of the videos.")
|
||||
parser.add_argument("--video_folder", type=str, default="", help="The video folder.")
|
||||
parser.add_argument(
|
||||
"--video_metadata_path", type=str, default=None, help="The path to the video dataset metadata (csv/jsonl)."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--video_path_column",
|
||||
type=str,
|
||||
default="video_path",
|
||||
help="The column contains the video path (an absolute path or a relative path w.r.t the video_folder).",
|
||||
)
|
||||
parser.add_argument("--saved_path", type=str, required=True, help="The save path to the output results (csv/jsonl).")
|
||||
parser.add_argument("--saved_freq", type=int, default=100, help="The frequency to save the output results.")
|
||||
parser.add_argument("--n_jobs", type=int, default=1, help="The number of concurrent processes.")
|
||||
|
||||
parser.add_argument(
|
||||
"--basic_metadata_path", type=str, default=None, help="The path to the basic metadata (csv/jsonl)."
|
||||
)
|
||||
parser.add_argument("--min_resolution", type=float, default=0, help="The resolution threshold.")
|
||||
parser.add_argument("--min_duration", type=float, default=-1, help="The minimum duration.")
|
||||
parser.add_argument("--max_duration", type=float, default=-1, help="The maximum duration.")
|
||||
parser.add_argument(
|
||||
"--asethetic_score_metadata_path", type=str, default=None, help="The path to the video quality metadata (csv/jsonl)."
|
||||
)
|
||||
parser.add_argument("--min_asethetic_score", type=float, default=4.0, help="The asethetic score threshold.")
|
||||
parser.add_argument(
|
||||
"--asethetic_score_siglip_metadata_path", type=str, default=None, help="The path to the video quality metadata (csv/jsonl)."
|
||||
)
|
||||
parser.add_argument("--min_asethetic_score_siglip", type=float, default=4.0, help="The asethetic score (SigLIP) threshold.")
|
||||
parser.add_argument(
|
||||
"--text_score_metadata_path", type=str, default=None, help="The path to the video text score metadata (csv/jsonl)."
|
||||
)
|
||||
parser.add_argument("--min_text_score", type=float, default=0.02, help="The text threshold.")
|
||||
|
||||
args = parser.parse_args()
|
||||
return args
|
||||
|
||||
|
||||
def main():
|
||||
args = parse_args()
|
||||
|
||||
if args.video_metadata_path.endswith(".csv"):
|
||||
video_metadata_df = pd.read_csv(args.video_metadata_path)
|
||||
elif args.video_metadata_path.endswith(".jsonl"):
|
||||
video_metadata_df = pd.read_json(args.video_metadata_path, lines=True)
|
||||
else:
|
||||
raise ValueError("The video_metadata_path must end with .csv or .jsonl.")
|
||||
video_path_list = video_metadata_df[args.video_path_column].tolist()
|
||||
|
||||
if not (args.saved_path.endswith(".csv") or args.saved_path.endswith(".jsonl")):
|
||||
raise ValueError("The saved_path must end with .csv or .jsonl.")
|
||||
|
||||
if os.path.exists(args.saved_path):
|
||||
if args.saved_path.endswith(".csv"):
|
||||
saved_metadata_df = pd.read_csv(args.saved_path)
|
||||
elif args.saved_path.endswith(".jsonl"):
|
||||
saved_metadata_df = pd.read_json(args.saved_path, lines=True)
|
||||
saved_video_path_list = saved_metadata_df[args.video_path_column].tolist()
|
||||
video_path_list = list(set(video_path_list).difference(set(saved_video_path_list)))
|
||||
logger.info(f"Resume from {args.saved_path}: {len(saved_video_path_list)} processed and {len(video_path_list)} to be processed.")
|
||||
|
||||
video_path_list = filter(
|
||||
video_path_list,
|
||||
basic_metadata_path=args.basic_metadata_path,
|
||||
min_resolution=args.min_resolution,
|
||||
min_duration=args.min_duration,
|
||||
max_duration=args.max_duration,
|
||||
asethetic_score_metadata_path=args.asethetic_score_metadata_path,
|
||||
min_asethetic_score=args.min_asethetic_score,
|
||||
asethetic_score_siglip_metadata_path=args.asethetic_score_siglip_metadata_path,
|
||||
min_asethetic_score_siglip=args.min_asethetic_score_siglip,
|
||||
text_score_metadata_path=args.text_score_metadata_path,
|
||||
min_text_score=args.min_text_score,
|
||||
)
|
||||
video_path_list = [os.path.join(args.video_folder, video_path) for video_path in video_path_list]
|
||||
# Sorting to guarantee the same result for each process.
|
||||
video_path_list = natsorted(video_path_list)
|
||||
|
||||
for i in tqdm(range(0, len(video_path_list), args.saved_freq)):
|
||||
result_list = Parallel(n_jobs=args.n_jobs)(
|
||||
delayed(compute_motion_score)(video_path) for video_path in tqdm(video_path_list[i: i + args.saved_freq])
|
||||
)
|
||||
result_list = [result for result in result_list if result is not None]
|
||||
if len(result_list) == 0:
|
||||
continue
|
||||
|
||||
result_df = pd.DataFrame(result_list)
|
||||
if args.saved_path.endswith(".csv"):
|
||||
header = False if os.path.exists(args.saved_path) else True
|
||||
result_df.to_csv(args.saved_path, header=header, index=False, mode="a")
|
||||
elif args.saved_path.endswith(".jsonl"):
|
||||
result_df.to_json(args.saved_path, orient="records", lines=True, mode="a", force_ascii=False)
|
||||
logger.info(f"Save result to {args.saved_path}.")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,214 @@
|
||||
import argparse
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
import easyocr
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
from accelerate import PartialState
|
||||
from accelerate.utils import gather_object
|
||||
from natsort import natsorted
|
||||
from tqdm import tqdm
|
||||
from torchvision.datasets.utils import download_url
|
||||
|
||||
from utils.logger import logger
|
||||
from utils.video_utils import extract_frames
|
||||
from utils.filter import filter
|
||||
|
||||
|
||||
def init_ocr_reader(root: str = "~/.cache/easyocr", device: str = "gpu"):
|
||||
root = os.path.expanduser(root)
|
||||
if not os.path.exists(root):
|
||||
os.makedirs(root)
|
||||
download_url(
|
||||
"https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/video_caption/easyocr/craft_mlt_25k.pth",
|
||||
root,
|
||||
filename="craft_mlt_25k.pth",
|
||||
md5="2f8227d2def4037cdb3b34389dcf9ec1",
|
||||
)
|
||||
ocr_reader = easyocr.Reader(
|
||||
lang_list=["en", "ch_sim"],
|
||||
gpu=device,
|
||||
recognizer=False,
|
||||
verbose=False,
|
||||
model_storage_directory=root,
|
||||
)
|
||||
|
||||
return ocr_reader
|
||||
|
||||
|
||||
def triangle_area(p1, p2, p3):
|
||||
"""Compute the triangle area according to its coordinates.
|
||||
"""
|
||||
x1, y1 = p1
|
||||
x2, y2 = p2
|
||||
x3, y3 = p3
|
||||
tri_area = 0.5 * np.abs(x1 * y2 + x2 * y3 + x3 * y1 - x2 * y1 - x3 * y2 - x1 * y3)
|
||||
return tri_area
|
||||
|
||||
|
||||
def compute_text_score(video_path, ocr_reader):
|
||||
_, images = extract_frames(video_path, sample_method="mid")
|
||||
images = [np.array(image) for image in images]
|
||||
|
||||
frame_ocr_area_ratios = []
|
||||
for image in images:
|
||||
# horizontal detected results and free-form detected
|
||||
horizontal_list, free_list = ocr_reader.detect(np.asarray(image))
|
||||
width, height = image.shape[0], image.shape[1]
|
||||
|
||||
total_area = width * height
|
||||
# rectangles
|
||||
rect_area = 0
|
||||
for xmin, xmax, ymin, ymax in horizontal_list[0]:
|
||||
if xmax < xmin or ymax < ymin:
|
||||
continue
|
||||
rect_area += (xmax - xmin) * (ymax - ymin)
|
||||
# free-form
|
||||
quad_area = 0
|
||||
try:
|
||||
for points in free_list[0]:
|
||||
triangle1 = points[:3]
|
||||
quad_area += triangle_area(*triangle1)
|
||||
triangle2 = points[3:] + [points[0]]
|
||||
quad_area += triangle_area(*triangle2)
|
||||
except:
|
||||
quad_area = 0
|
||||
text_area = rect_area + quad_area
|
||||
|
||||
frame_ocr_area_ratios.append(text_area / total_area)
|
||||
|
||||
video_meta_info = {
|
||||
"video_path": Path(video_path).name,
|
||||
"text_score": round(np.mean(frame_ocr_area_ratios), 5),
|
||||
}
|
||||
|
||||
return video_meta_info
|
||||
|
||||
|
||||
def parse_args():
|
||||
parser = argparse.ArgumentParser(description="Compute the text score of the middle frame in the videos.")
|
||||
parser.add_argument("--video_folder", type=str, default="", help="The video folder.")
|
||||
parser.add_argument(
|
||||
"--video_metadata_path", type=str, default=None, help="The path to the video dataset metadata (csv/jsonl)."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--video_path_column",
|
||||
type=str,
|
||||
default="video_path",
|
||||
help="The column contains the video path (an absolute path or a relative path w.r.t the video_folder).",
|
||||
)
|
||||
parser.add_argument("--saved_path", type=str, required=True, help="The save path to the output results (csv/jsonl).")
|
||||
parser.add_argument("--saved_freq", type=int, default=100, help="The frequency to save the output results.")
|
||||
|
||||
parser.add_argument(
|
||||
"--basic_metadata_path", type=str, default=None, help="The path to the basic metadata (csv/jsonl)."
|
||||
)
|
||||
parser.add_argument("--min_resolution", type=float, default=0, help="The resolution threshold.")
|
||||
parser.add_argument("--min_duration", type=float, default=-1, help="The minimum duration.")
|
||||
parser.add_argument("--max_duration", type=float, default=-1, help="The maximum duration.")
|
||||
parser.add_argument(
|
||||
"--asethetic_score_metadata_path", type=str, default=None, help="The path to the video quality metadata (csv/jsonl)."
|
||||
)
|
||||
parser.add_argument("--min_asethetic_score", type=float, default=4.0, help="The asethetic score threshold.")
|
||||
parser.add_argument(
|
||||
"--asethetic_score_siglip_metadata_path", type=str, default=None, help="The path to the video quality metadata (csv/jsonl)."
|
||||
)
|
||||
parser.add_argument("--min_asethetic_score_siglip", type=float, default=4.0, help="The asethetic score (SigLIP) threshold.")
|
||||
parser.add_argument(
|
||||
"--motion_score_metadata_path", type=str, default=None, help="The path to the video motion score metadata (csv/jsonl)."
|
||||
)
|
||||
parser.add_argument("--min_motion_score", type=float, default=2, help="The motion threshold.")
|
||||
|
||||
args = parser.parse_args()
|
||||
return args
|
||||
|
||||
|
||||
def main():
|
||||
args = parse_args()
|
||||
|
||||
if args.video_metadata_path.endswith(".csv"):
|
||||
video_metadata_df = pd.read_csv(args.video_metadata_path)
|
||||
elif args.video_metadata_path.endswith(".jsonl"):
|
||||
video_metadata_df = pd.read_json(args.video_metadata_path, lines=True)
|
||||
else:
|
||||
raise ValueError("The video_metadata_path must end with .csv or .jsonl.")
|
||||
video_path_list = video_metadata_df[args.video_path_column].tolist()
|
||||
|
||||
if not (args.saved_path.endswith(".csv") or args.saved_path.endswith(".jsonl")):
|
||||
raise ValueError("The saved_path must end with .csv or .jsonl.")
|
||||
|
||||
if os.path.exists(args.saved_path):
|
||||
if args.saved_path.endswith(".csv"):
|
||||
saved_metadata_df = pd.read_csv(args.saved_path)
|
||||
elif args.saved_path.endswith(".jsonl"):
|
||||
saved_metadata_df = pd.read_json(args.saved_path, lines=True)
|
||||
saved_video_path_list = saved_metadata_df[args.video_path_column].tolist()
|
||||
video_path_list = list(set(video_path_list).difference(set(saved_video_path_list)))
|
||||
logger.info(f"Resume from {args.saved_path}: {len(saved_video_path_list)} processed and {len(video_path_list)} to be processed.")
|
||||
|
||||
video_path_list = filter(
|
||||
video_path_list,
|
||||
basic_metadata_path=args.basic_metadata_path,
|
||||
min_resolution=args.min_resolution,
|
||||
min_duration=args.min_duration,
|
||||
max_duration=args.max_duration,
|
||||
asethetic_score_metadata_path=args.asethetic_score_metadata_path,
|
||||
min_asethetic_score=args.min_asethetic_score,
|
||||
asethetic_score_siglip_metadata_path=args.asethetic_score_siglip_metadata_path,
|
||||
min_asethetic_score_siglip=args.min_asethetic_score_siglip,
|
||||
motion_score_metadata_path=args.motion_score_metadata_path,
|
||||
min_motion_score=args.min_motion_score,
|
||||
)
|
||||
video_path_list = [os.path.join(args.video_folder, video_path) for video_path in video_path_list]
|
||||
# Sorting to guarantee the same result for each process.
|
||||
video_path_list = natsorted(video_path_list)
|
||||
|
||||
state = PartialState()
|
||||
if state.is_main_process:
|
||||
# Check if the model is downloaded in the main process.
|
||||
ocr_reader = init_ocr_reader(device="cpu")
|
||||
state.wait_for_everyone()
|
||||
ocr_reader = init_ocr_reader(device=state.device)
|
||||
|
||||
index = len(video_path_list) - len(video_path_list) % state.num_processes
|
||||
# Avoid the NCCL timeout in the final gather operation.
|
||||
logger.info(f"Drop {len(video_path_list) % state.num_processes} videos to ensure each process handles the same number of videos.")
|
||||
video_path_list = video_path_list[:index]
|
||||
logger.info(f"{len(video_path_list)} videos are to be processed.")
|
||||
|
||||
result_list = []
|
||||
with state.split_between_processes(video_path_list) as splitted_video_path_list:
|
||||
for i, video_path in enumerate(tqdm(splitted_video_path_list)):
|
||||
try:
|
||||
video_meta_info = compute_text_score(video_path, ocr_reader)
|
||||
result_list.append(video_meta_info)
|
||||
except Exception as e:
|
||||
logger.warning(f"Compute text score for video {video_path} with error: {e}.")
|
||||
if i != 0 and i % args.saved_freq == 0:
|
||||
state.wait_for_everyone()
|
||||
gathered_result_list = gather_object(result_list)
|
||||
if state.is_main_process and len(gathered_result_list) != 0:
|
||||
result_df = pd.DataFrame(gathered_result_list)
|
||||
if args.saved_path.endswith(".csv"):
|
||||
header = False if os.path.exists(args.saved_path) else True
|
||||
result_df.to_csv(args.saved_path, header=header, index=False, mode="a")
|
||||
elif args.saved_path.endswith(".jsonl"):
|
||||
result_df.to_json(args.saved_path, orient="records", lines=True, mode="a", force_ascii=False)
|
||||
logger.info(f"Save result to {args.saved_path}.")
|
||||
result_list = []
|
||||
|
||||
state.wait_for_everyone()
|
||||
gathered_result_list = gather_object(result_list)
|
||||
if state.is_main_process and len(gathered_result_list) != 0:
|
||||
result_df = pd.DataFrame(gathered_result_list)
|
||||
if args.saved_path.endswith(".csv"):
|
||||
header = False if os.path.exists(args.saved_path) else True
|
||||
result_df.to_csv(args.saved_path, header=header, index=False, mode="a")
|
||||
elif args.saved_path.endswith(".jsonl"):
|
||||
result_df.to_json(args.saved_path, orient="records", lines=True, mode="a", force_ascii=False)
|
||||
logger.info(f"Save the final result to {args.saved_path}.")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,201 @@
|
||||
import argparse
|
||||
import os
|
||||
|
||||
import pandas as pd
|
||||
from accelerate import PartialState
|
||||
from accelerate.utils import gather_object
|
||||
from natsort import index_natsorted
|
||||
from tqdm import tqdm
|
||||
from torch.utils.data import DataLoader
|
||||
|
||||
import utils.image_evaluator as image_evaluator
|
||||
import utils.video_evaluator as video_evaluator
|
||||
from utils.logger import logger
|
||||
from utils.video_dataset import VideoDataset, collate_fn
|
||||
|
||||
|
||||
def parse_args():
|
||||
parser = argparse.ArgumentParser(description="Compute scores of uniform sampled frames from videos.")
|
||||
parser.add_argument(
|
||||
"--video_metadata_path", type=str, default=None, help="The path to the video dataset metadata (csv/jsonl)."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--video_path_column",
|
||||
type=str,
|
||||
default="video_path",
|
||||
help="The column contains the video path (an absolute path or a relative path w.r.t the video_folder).",
|
||||
)
|
||||
parser.add_argument("--video_folder", type=str, default="", help="The video folder.")
|
||||
parser.add_argument(
|
||||
"--caption_column",
|
||||
type=str,
|
||||
default=None,
|
||||
help="The column contains the caption.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--frame_sample_method",
|
||||
type=str,
|
||||
choices=["mid", "uniform", "image"],
|
||||
default="uniform",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--num_sampled_frames",
|
||||
type=int,
|
||||
default=8,
|
||||
help="num_sampled_frames",
|
||||
)
|
||||
parser.add_argument("--metrics", nargs="+", type=str, required=True, help="The evaluation metric(s) for generated images.")
|
||||
parser.add_argument(
|
||||
"--batch_size",
|
||||
type=int,
|
||||
default=10,
|
||||
required=False,
|
||||
help="The batch size for the video dataset.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--num_workers",
|
||||
type=int,
|
||||
default=4,
|
||||
required=False,
|
||||
help="The number of workers for the video dataset.",
|
||||
)
|
||||
parser.add_argument("--saved_path", type=str, required=True, help="The save path to the output results (csv/jsonl).")
|
||||
parser.add_argument("--saved_freq", type=int, default=1000, help="The frequency to save the output results.")
|
||||
|
||||
args = parser.parse_args()
|
||||
return args
|
||||
|
||||
|
||||
def main():
|
||||
args = parse_args()
|
||||
|
||||
if args.video_metadata_path.endswith(".csv"):
|
||||
video_metadata_df = pd.read_csv(args.video_metadata_path)
|
||||
elif args.video_metadata_path.endswith(".jsonl"):
|
||||
video_metadata_df = pd.read_json(args.video_metadata_path, lines=True)
|
||||
else:
|
||||
raise ValueError("The video_metadata_path must end with .csv or .jsonl.")
|
||||
|
||||
if not (args.saved_path.endswith(".csv") or args.saved_path.endswith(".jsonl")):
|
||||
raise ValueError("The saved_path must end with .csv or .jsonl.")
|
||||
|
||||
if os.path.exists(args.saved_path):
|
||||
if args.saved_path.endswith(".csv"):
|
||||
saved_metadata_df = pd.read_csv(args.saved_path)
|
||||
elif args.saved_path.endswith(".jsonl"):
|
||||
saved_metadata_df = pd.read_json(args.saved_path, lines=True)
|
||||
|
||||
# Filter out the unprocessed video-caption pairs by setting the indicator=True.
|
||||
merged_df = video_metadata_df.merge(saved_metadata_df, on="video_path", how="outer", indicator=True)
|
||||
video_metadata_df = merged_df[merged_df["_merge"] == "left_only"]
|
||||
# Sorting to guarantee the same result for each process.
|
||||
video_metadata_df = video_metadata_df.iloc[index_natsorted(video_metadata_df["video_path"])].reset_index(drop=True)
|
||||
if args.caption_column is None:
|
||||
video_metadata_df = video_metadata_df[[args.video_path_column]]
|
||||
else:
|
||||
video_metadata_df = video_metadata_df[[args.video_path_column, args.caption_column + "_x"]]
|
||||
video_metadata_df.rename(columns={args.caption_column + "_x": args.caption_column}, inplace=True)
|
||||
logger.info(f"Resume from {args.saved_path}: {len(saved_metadata_df)} processed and {len(video_metadata_df)} to be processed.")
|
||||
|
||||
state = PartialState()
|
||||
metric_fns = []
|
||||
for metric in args.metrics:
|
||||
if hasattr(image_evaluator, metric): # frame-wise
|
||||
if state.is_main_process:
|
||||
logger.info("Initializing frame-wise evaluator metrics...")
|
||||
# Check if the model is downloaded in the main process.
|
||||
getattr(image_evaluator, metric)(device="cpu")
|
||||
state.wait_for_everyone()
|
||||
metric_fns.append(getattr(image_evaluator, metric)(device=state.device))
|
||||
else: # video-wise
|
||||
if state.is_main_process:
|
||||
logger.info("Initializing video-wise evaluator metrics...")
|
||||
# Check if the model is downloaded in the main process.
|
||||
getattr(video_evaluator, metric)(device="cpu")
|
||||
state.wait_for_everyone()
|
||||
metric_fns.append(getattr(video_evaluator, metric)(device=state.device))
|
||||
|
||||
result_dict = {args.video_path_column: [], "sample_frame_idx": []}
|
||||
for metric in metric_fns:
|
||||
result_dict[str(metric)] = []
|
||||
if args.caption_column is not None:
|
||||
result_dict[args.caption_column] = []
|
||||
|
||||
if args.frame_sample_method == "image":
|
||||
logger.warning("Set args.num_sampled_frames to 1 since args.frame_sample_method is image.")
|
||||
args.num_sampled_frames = 1
|
||||
|
||||
index = len(video_metadata_df) - len(video_metadata_df) % state.num_processes
|
||||
# Avoid the NCCL timeout in the final gather operation.
|
||||
logger.info(f"Drop {len(video_metadata_df) % state.num_processes} videos to ensure each process handles the same number of videos.")
|
||||
video_metadata_df = video_metadata_df.iloc[:index]
|
||||
logger.info(f"{len(video_metadata_df)} videos are to be processed.")
|
||||
|
||||
video_metadata_list = video_metadata_df.to_dict(orient='list')
|
||||
with state.split_between_processes(video_metadata_list) as splitted_video_metadata:
|
||||
video_dataset = VideoDataset(
|
||||
dataset_inputs=splitted_video_metadata,
|
||||
video_folder=args.video_folder,
|
||||
text_column=args.caption_column,
|
||||
sample_method=args.frame_sample_method,
|
||||
num_sampled_frames=args.num_sampled_frames
|
||||
)
|
||||
video_loader = DataLoader(video_dataset, batch_size=args.batch_size, num_workers=args.num_workers, collate_fn=collate_fn)
|
||||
|
||||
for idx, batch in enumerate(tqdm(video_loader)):
|
||||
if len(batch) > 0:
|
||||
batch_video_path = batch["path"]
|
||||
result_dict["sample_frame_idx"].extend(batch["sampled_frame_idx"])
|
||||
batch_frame = batch["sampled_frame"] # [batch_size, num_sampled_frames, H, W, C]
|
||||
batch_caption = None
|
||||
if args.caption_column is not None:
|
||||
batch_caption = batch["text"]
|
||||
result_dict["caption"].extend(batch_caption)
|
||||
# Compute the quality.
|
||||
for i, metric in enumerate(args.metrics):
|
||||
quality_scores = metric_fns[i](batch_frame, batch_caption)
|
||||
if isinstance(quality_scores[0], list): # frame-wise
|
||||
quality_scores = [
|
||||
[round(score, 5) for score in inner_list]
|
||||
for inner_list in quality_scores
|
||||
]
|
||||
else: # video-wise
|
||||
quality_scores = [round(score, 5) for score in quality_scores]
|
||||
result_dict[str(metric_fns[i])].extend(quality_scores)
|
||||
|
||||
if args.video_folder == "":
|
||||
saved_video_path_list = batch_video_path
|
||||
else:
|
||||
saved_video_path_list = [os.path.relpath(video_path, args.video_folder) for video_path in batch_video_path]
|
||||
result_dict[args.video_path_column].extend(saved_video_path_list)
|
||||
|
||||
# Save the metadata in the main process every saved_freq.
|
||||
if (idx != 0) and (idx % args.saved_freq == 0):
|
||||
state.wait_for_everyone()
|
||||
gathered_result_dict = {k: gather_object(v) for k, v in result_dict.items()}
|
||||
if state.is_main_process and len(gathered_result_dict[args.video_path_column]) != 0:
|
||||
result_df = pd.DataFrame(gathered_result_dict)
|
||||
if args.saved_path.endswith(".csv"):
|
||||
header = False if os.path.exists(args.saved_path) else True
|
||||
result_df.to_csv(args.saved_path, header=header, index=False, mode="a")
|
||||
elif args.saved_path.endswith(".jsonl"):
|
||||
result_df.to_json(args.saved_path, orient="records", lines=True, mode="a", force_ascii=False)
|
||||
logger.info(f"Save result to {args.saved_path}.")
|
||||
for k in result_dict.keys():
|
||||
result_dict[k] = []
|
||||
|
||||
# Wait for all processes to finish and gather the final result.
|
||||
state.wait_for_everyone()
|
||||
gathered_result_dict = {k: gather_object(v) for k, v in result_dict.items()}
|
||||
# Save the metadata in the main process.
|
||||
if state.is_main_process and len(gathered_result_dict[args.video_path_column]) != 0:
|
||||
result_df = pd.DataFrame(gathered_result_dict)
|
||||
if args.saved_path.endswith(".csv"):
|
||||
header = False if os.path.exists(args.saved_path) else True
|
||||
result_df.to_csv(args.saved_path, header=header, index=False, mode="a")
|
||||
elif args.saved_path.endswith(".jsonl"):
|
||||
result_df.to_json(args.saved_path, orient="records", lines=True, mode="a", force_ascii=False)
|
||||
logger.info(f"Save the final result to {args.saved_path}.")
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,97 @@
|
||||
import argparse
|
||||
import os
|
||||
from copy import deepcopy
|
||||
from pathlib import Path
|
||||
from multiprocessing import Pool
|
||||
|
||||
import pandas as pd
|
||||
from scenedetect import open_video, SceneManager
|
||||
from scenedetect.detectors import ContentDetector
|
||||
from tqdm import tqdm
|
||||
|
||||
from utils.logger import logger
|
||||
|
||||
|
||||
def cutscene_detection_star(args):
|
||||
return cutscene_detection(*args)
|
||||
|
||||
|
||||
def cutscene_detection(video_path, saved_path, cutscene_threshold=27, min_scene_len=15):
|
||||
try:
|
||||
if os.path.exists(saved_path):
|
||||
logger.info(f"{video_path} has been processed.")
|
||||
return
|
||||
# Use PyAV as the backend to avoid (to some exent) containing the last frame of the previous scene.
|
||||
# https://github.com/Breakthrough/PySceneDetect/issues/279#issuecomment-2152596761.
|
||||
video = open_video(video_path, backend="pyav")
|
||||
frame_rate, frame_size = video.frame_rate, video.frame_size
|
||||
duration = deepcopy(video.duration)
|
||||
|
||||
frame_points, frame_timecode = [], {}
|
||||
scene_manager = SceneManager()
|
||||
scene_manager.add_detector(
|
||||
# [ContentDetector, ThresholdDetector, AdaptiveDetector]
|
||||
ContentDetector(threshold=cutscene_threshold, min_scene_len=min_scene_len)
|
||||
)
|
||||
scene_manager.detect_scenes(video, show_progress=False)
|
||||
scene_list = scene_manager.get_scene_list()
|
||||
for scene in scene_list:
|
||||
for frame_time_code in scene:
|
||||
frame_index = frame_time_code.get_frames()
|
||||
if frame_index not in frame_points:
|
||||
frame_points.append(frame_index)
|
||||
frame_timecode[frame_index] = frame_time_code
|
||||
|
||||
del video, scene_manager
|
||||
|
||||
frame_points = sorted(frame_points)
|
||||
output_scene_list = []
|
||||
for idx in range(len(frame_points) - 1):
|
||||
output_scene_list.append((frame_timecode[frame_points[idx]], frame_timecode[frame_points[idx+1]]))
|
||||
|
||||
timecode_list = [(frame_timecode_tuple[0].get_timecode(), frame_timecode_tuple[1].get_timecode()) for frame_timecode_tuple in output_scene_list]
|
||||
meta_scene = [{
|
||||
"video_path": Path(video_path).name,
|
||||
"timecode_list": timecode_list,
|
||||
"fram_rate": frame_rate,
|
||||
"frame_size": frame_size,
|
||||
"duration": str(duration) # __repr__
|
||||
}]
|
||||
pd.DataFrame(meta_scene).to_json(saved_path, orient="records", lines=True)
|
||||
except Exception as e:
|
||||
logger.warning(f"Cutscene detection with {video_path} failed. Error is: {e}.")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser(description="Cutscene Detection")
|
||||
parser.add_argument(
|
||||
"--video_metadata_path", type=str, required=True, help="The path to the video dataset metadata (csv/jsonl)."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--video_path_column",
|
||||
type=str,
|
||||
default="video_path",
|
||||
help="The column contains the video path (an absolute path or a relative path w.r.t the video_folder).",
|
||||
)
|
||||
parser.add_argument("--video_folder", type=str, default="", help="The video folder.")
|
||||
parser.add_argument("--saved_folder", type=str, required=True, help="The save path to the output results (csv/jsonl).")
|
||||
parser.add_argument("--n_jobs", type=int, default=1, help="The number of processes.")
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
metadata_df = pd.read_json(args.video_metadata_path, lines=True)
|
||||
video_path_list = metadata_df[args.video_path_column].tolist()
|
||||
video_path_list = [os.path.join(args.video_folder, video_path) for video_path in video_path_list]
|
||||
|
||||
if not os.path.exists(args.saved_folder):
|
||||
os.makedirs(args.saved_folder, exist_ok=True)
|
||||
# The glob can be slow when there are many small jsonl files.
|
||||
saved_path_list = [os.path.join(args.saved_folder, Path(video_path).stem + ".jsonl") for video_path in video_path_list]
|
||||
args_list = [
|
||||
(video_path, saved_path)
|
||||
for video_path, saved_path in zip(video_path_list, saved_path_list)
|
||||
]
|
||||
# Since the length of the video is not uniform, the gather operation is not performed.
|
||||
# We need to run easyanimate/video_caption/utils/gather_jsonl.py after the program finised.
|
||||
with Pool(args.n_jobs) as pool:
|
||||
results = list(tqdm(pool.imap(cutscene_detection_star, args_list), total=len(video_path_list)))
|
||||
@@ -0,0 +1,88 @@
|
||||
import argparse
|
||||
import os
|
||||
|
||||
import pandas as pd
|
||||
from natsort import natsorted
|
||||
|
||||
from utils.logger import logger
|
||||
from utils.filter import filter
|
||||
|
||||
|
||||
def parse_args():
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument(
|
||||
"--caption_metadata_path", type=str, default=None, help="The path to the video quality metadata (csv/jsonl)."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--video_path_column",
|
||||
type=str,
|
||||
default="video_path",
|
||||
help="The column contains the video path (an absolute path or a relative path w.r.t the video_folder).",
|
||||
)
|
||||
parser.add_argument("--video_folder", type=str, default="", help="The video folder.")
|
||||
parser.add_argument(
|
||||
"--basic_metadata_path", type=str, default=None, help="The path to the basic metadata (csv/jsonl)."
|
||||
)
|
||||
parser.add_argument("--min_resolution", type=float, default=720*1280, help="The resolution threshold.")
|
||||
parser.add_argument("--min_duration", type=float, default=-1, help="The minimum duration.")
|
||||
parser.add_argument("--max_duration", type=float, default=-1, help="The maximum duration.")
|
||||
parser.add_argument(
|
||||
"--asethetic_score_metadata_path", type=str, default=None, help="The path to the video quality metadata (csv/jsonl)."
|
||||
)
|
||||
parser.add_argument("--min_asethetic_score", type=float, default=4.0, help="The asethetic score threshold.")
|
||||
parser.add_argument(
|
||||
"--asethetic_score_siglip_metadata_path", type=str, default=None, help="The path to the video quality (SigLIP) metadata (csv/jsonl)."
|
||||
)
|
||||
parser.add_argument("--min_asethetic_score_siglip", type=float, default=4.0, help="The asethetic score (SigLIP) threshold.")
|
||||
parser.add_argument(
|
||||
"--text_score_metadata_path", type=str, default=None, help="The path to the video text score metadata (csv/jsonl)."
|
||||
)
|
||||
parser.add_argument("--min_text_score", type=float, default=0.02, help="The text threshold.")
|
||||
parser.add_argument(
|
||||
"--motion_score_metadata_path", type=str, default=None, help="The path to the video motion score metadata (csv/jsonl)."
|
||||
)
|
||||
parser.add_argument("--min_motion_score", type=float, default=2, help="The motion threshold.")
|
||||
parser.add_argument(
|
||||
"--videoclipxl_score_metadata_path", type=str, default=None, help="The path to the video-caption VideoCLIPXL score metadata (csv/jsonl)."
|
||||
)
|
||||
parser.add_argument("--min_videoclipxl_score", type=float, default=0.20, help="The VideoCLIPXL score threshold.")
|
||||
parser.add_argument("--saved_path", type=str, required=True)
|
||||
|
||||
args = parser.parse_args()
|
||||
return args
|
||||
|
||||
|
||||
def main():
|
||||
args = parse_args()
|
||||
|
||||
raw_caption_df = pd.read_json(args.caption_metadata_path, lines=True)
|
||||
video_path_list = raw_caption_df[args.video_path_column].to_list()
|
||||
filtered_video_path_list = filter(
|
||||
video_path_list,
|
||||
basic_metadata_path=args.basic_metadata_path,
|
||||
min_resolution=args.min_resolution,
|
||||
min_duration=args.min_duration,
|
||||
max_duration=args.max_duration,
|
||||
asethetic_score_metadata_path=args.asethetic_score_metadata_path,
|
||||
min_asethetic_score=args.min_asethetic_score,
|
||||
asethetic_score_siglip_metadata_path=args.asethetic_score_siglip_metadata_path,
|
||||
min_asethetic_score_siglip=args.min_asethetic_score_siglip,
|
||||
text_score_metadata_path=args.text_score_metadata_path,
|
||||
min_text_score=args.min_text_score,
|
||||
motion_score_metadata_path=args.motion_score_metadata_path,
|
||||
min_motion_score=args.min_motion_score,
|
||||
videoclipxl_score_metadata_path=args.videoclipxl_score_metadata_path,
|
||||
min_videoclipxl_score=args.min_videoclipxl_score,
|
||||
video_path_column=args.video_path_column
|
||||
)
|
||||
filtered_video_path_list = natsorted(filtered_video_path_list)
|
||||
filtered_caption_df = raw_caption_df[raw_caption_df[args.video_path_column].isin(filtered_video_path_list)]
|
||||
train_df = filtered_caption_df.rename(columns={"video_path": "file_path", "caption": "text"})
|
||||
train_df["file_path"] = train_df["file_path"].map(lambda x: os.path.join(args.video_folder, x))
|
||||
train_df["type"] = "video"
|
||||
train_df.to_json(args.saved_path, orient="records", force_ascii=False, indent=2)
|
||||
logger.info(f"The final train file with {len(train_df)} videos are saved to {args.saved_path}.")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,114 @@
|
||||
"""Modified from https://github.com/JaidedAI/EasyOCR/blob/803b907/easyocr/detection.py.
|
||||
1. Disable DataParallel.
|
||||
"""
|
||||
import torch
|
||||
import torch.backends.cudnn as cudnn
|
||||
from torch.autograd import Variable
|
||||
from PIL import Image
|
||||
from collections import OrderedDict
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
from .craft_utils import getDetBoxes, adjustResultCoordinates
|
||||
from .imgproc import resize_aspect_ratio, normalizeMeanVariance
|
||||
from .craft import CRAFT
|
||||
|
||||
def copyStateDict(state_dict):
|
||||
if list(state_dict.keys())[0].startswith("module"):
|
||||
start_idx = 1
|
||||
else:
|
||||
start_idx = 0
|
||||
new_state_dict = OrderedDict()
|
||||
for k, v in state_dict.items():
|
||||
name = ".".join(k.split(".")[start_idx:])
|
||||
new_state_dict[name] = v
|
||||
return new_state_dict
|
||||
|
||||
def test_net(canvas_size, mag_ratio, net, image, text_threshold, link_threshold, low_text, poly, device, estimate_num_chars=False):
|
||||
if isinstance(image, np.ndarray) and len(image.shape) == 4: # image is batch of np arrays
|
||||
image_arrs = image
|
||||
else: # image is single numpy array
|
||||
image_arrs = [image]
|
||||
|
||||
img_resized_list = []
|
||||
# resize
|
||||
for img in image_arrs:
|
||||
img_resized, target_ratio, size_heatmap = resize_aspect_ratio(img, canvas_size,
|
||||
interpolation=cv2.INTER_LINEAR,
|
||||
mag_ratio=mag_ratio)
|
||||
img_resized_list.append(img_resized)
|
||||
ratio_h = ratio_w = 1 / target_ratio
|
||||
# preprocessing
|
||||
x = [np.transpose(normalizeMeanVariance(n_img), (2, 0, 1))
|
||||
for n_img in img_resized_list]
|
||||
x = torch.from_numpy(np.array(x))
|
||||
x = x.to(device)
|
||||
|
||||
# forward pass
|
||||
with torch.no_grad():
|
||||
y, feature = net(x)
|
||||
|
||||
boxes_list, polys_list = [], []
|
||||
for out in y:
|
||||
# make score and link map
|
||||
score_text = out[:, :, 0].cpu().data.numpy()
|
||||
score_link = out[:, :, 1].cpu().data.numpy()
|
||||
|
||||
# Post-processing
|
||||
boxes, polys, mapper = getDetBoxes(
|
||||
score_text, score_link, text_threshold, link_threshold, low_text, poly, estimate_num_chars)
|
||||
|
||||
# coordinate adjustment
|
||||
boxes = adjustResultCoordinates(boxes, ratio_w, ratio_h)
|
||||
polys = adjustResultCoordinates(polys, ratio_w, ratio_h)
|
||||
if estimate_num_chars:
|
||||
boxes = list(boxes)
|
||||
polys = list(polys)
|
||||
for k in range(len(polys)):
|
||||
if estimate_num_chars:
|
||||
boxes[k] = (boxes[k], mapper[k])
|
||||
if polys[k] is None:
|
||||
polys[k] = boxes[k]
|
||||
boxes_list.append(boxes)
|
||||
polys_list.append(polys)
|
||||
|
||||
return boxes_list, polys_list
|
||||
|
||||
def get_detector(trained_model, device='cpu', quantize=True, cudnn_benchmark=False):
|
||||
net = CRAFT()
|
||||
|
||||
if device == 'cpu':
|
||||
net.load_state_dict(copyStateDict(torch.load(trained_model, map_location=device)))
|
||||
if quantize:
|
||||
try:
|
||||
torch.quantization.quantize_dynamic(net, dtype=torch.qint8, inplace=True)
|
||||
except:
|
||||
pass
|
||||
else:
|
||||
net.load_state_dict(copyStateDict(torch.load(trained_model, map_location=device)))
|
||||
# net = torch.nn.DataParallel(net).to(device)
|
||||
net = net.to(device)
|
||||
cudnn.benchmark = cudnn_benchmark
|
||||
|
||||
net.eval()
|
||||
return net
|
||||
|
||||
def get_textbox(detector, image, canvas_size, mag_ratio, text_threshold, link_threshold, low_text, poly, device, optimal_num_chars=None, **kwargs):
|
||||
result = []
|
||||
estimate_num_chars = optimal_num_chars is not None
|
||||
bboxes_list, polys_list = test_net(canvas_size, mag_ratio, detector,
|
||||
image, text_threshold,
|
||||
link_threshold, low_text, poly,
|
||||
device, estimate_num_chars)
|
||||
if estimate_num_chars:
|
||||
polys_list = [[p for p, _ in sorted(polys, key=lambda x: abs(optimal_num_chars - x[1]))]
|
||||
for polys in polys_list]
|
||||
|
||||
for polys in polys_list:
|
||||
single_img_result = []
|
||||
for i, box in enumerate(polys):
|
||||
poly = np.array(box).astype(np.int32).reshape((-1))
|
||||
single_img_result.append(poly)
|
||||
result.append(single_img_result)
|
||||
|
||||
return result
|
||||
@@ -0,0 +1,42 @@
|
||||
# Modified from https://github.com/NVlabs/VILA/blob/1c88211/llava/model/multimodal_encoder/siglip_encoder.py
|
||||
# 1. Support transformers >= 4.36.2.
|
||||
import torch
|
||||
import transformers
|
||||
from packaging import version
|
||||
from transformers import AutoConfig, AutoModel, PretrainedConfig
|
||||
|
||||
from llava.model.multimodal_encoder.vision_encoder import VisionTower, VisionTowerS2
|
||||
|
||||
if version.parse(transformers.__version__) > version.parse("4.36.2"):
|
||||
from transformers import SiglipImageProcessor, SiglipVisionConfig, SiglipVisionModel
|
||||
else:
|
||||
from .siglip import SiglipImageProcessor, SiglipVisionConfig, SiglipVisionModel
|
||||
|
||||
|
||||
class SiglipVisionTower(VisionTower):
|
||||
def __init__(self, model_name_or_path: str, config: PretrainedConfig, state_dict=None):
|
||||
super().__init__(model_name_or_path, config)
|
||||
self.image_processor = SiglipImageProcessor.from_pretrained(model_name_or_path)
|
||||
self.vision_tower = SiglipVisionModel.from_pretrained(
|
||||
# TODO(ligeng): why pass config here leading to errors?
|
||||
model_name_or_path, torch_dtype=eval(config.model_dtype), state_dict=state_dict
|
||||
)
|
||||
self.is_loaded = True
|
||||
|
||||
|
||||
class SiglipVisionTowerS2(VisionTowerS2):
|
||||
def __init__(self, model_name_or_path: str, config: PretrainedConfig):
|
||||
super().__init__(model_name_or_path, config)
|
||||
self.image_processor = SiglipImageProcessor.from_pretrained(model_name_or_path)
|
||||
self.vision_tower = SiglipVisionModel.from_pretrained(
|
||||
model_name_or_path, torch_dtype=eval(config.model_dtype)
|
||||
)
|
||||
|
||||
# Make sure it crops/resizes the image to the largest scale in self.scales to maintain high-res information
|
||||
self.image_processor.size['height'] = self.image_processor.size['width'] = self.scales[-1]
|
||||
|
||||
self.is_loaded = True
|
||||
|
||||
if version.parse(transformers.__version__) <= version.parse("4.36.2"):
|
||||
AutoConfig.register("siglip_vision_model", SiglipVisionConfig)
|
||||
AutoModel.register(SiglipVisionConfig, SiglipVisionModel)
|
||||
@@ -0,0 +1,9 @@
|
||||
I will upload some brief prompt words to be used for AI-generated videos. Please expand these brief prompt words into a more detailed description to enhance the quality of the generated videos. The detailed description should include the main subject (person, object, animal, or none) actions and their attributes or status sequence, the background (the objects, location, weather, and time), the view shot and camera movement.
|
||||
The final detailed description must not exceed 200 words. Output with the following json format:
|
||||
{"detailed description": "your detailed description here"}
|
||||
|
||||
Here is an example:
|
||||
brief prompt words: "A stylish woman in a black leather jacket, red dress, and boots walks confidently down a damp Tokyo street."
|
||||
{"detailed description": "A stylish woman walks down a Tokyo street filled with warm glowing neon and animated city signage. She wears a black leather jacket, a long red dress, and black boots, and carries a black purse. She wears sunglasses and red lipstick. She walks confidently and casually. The street is damp and reflective, creating a mirror effect of the colorful lights. Many pedestrians walk about."}
|
||||
|
||||
Here are the brief prompt words:
|
||||
@@ -0,0 +1,9 @@
|
||||
Please rewrite the video description to be useful for AI to re-generate the video, according to the following requirements
|
||||
1. Do not start with something similar to 'The video/scene/frame shows' or "In this video/scene/frame".
|
||||
2. Remove the subjective content deviates from describing the visual content of the video. For instance, a sentence like "It gives a feeling of ease and tranquility and makes people feel comfortable" is considered subjective.
|
||||
3. Remove the non-existent description that does not in the visual content of the video, For instance, a sentence like "There is no visible detail that could be used to identify the individual beyond what is shown." is considered as the non-existent description.
|
||||
4. Here are some examples of good descriptions: 1) A stylish woman walks down a Tokyo street filled with warm glowing neon and animated city signage. She wears a black leather jacket, a long red dress, and black boots, and carries a black purse. She wears sunglasses and red lipstick. She walks confidently and casually. The street is damp and reflective, creating a mirror effect of the colorful lights. Many pedestrians walk about. 2) A large orange octopus is seen resting on the bottom of the ocean floor, blending in with the sandy and rocky terrain. Its tentacles are spread out around its body, and its eyes are closed. The octopus is unaware of a king crab that is crawling towards it from behind a rock, its claws raised and ready to attack. The crab is brown and spiny, with long legs and antennae. The scene is captured from a wide angle, showing the vastness and depth of the ocean. The water is clear and blue, with rays of sunlight filtering through. The shot is sharp and crisp, with a high dynamic range. The octopus and the crab are in focus, while the background is slightly blurred, creating a depth of field effect.
|
||||
5. Output with the following json format:
|
||||
{"rewritten description": "your rewritten description here"}
|
||||
|
||||
Here is the video description:
|
||||
@@ -0,0 +1,9 @@
|
||||
pandas>=2.0.0
|
||||
easyocr==1.7.1
|
||||
git+https://github.com/openai/CLIP.git
|
||||
natsort
|
||||
joblib
|
||||
scenedetect
|
||||
av
|
||||
# https://github.com/NVlabs/VILA/issues/78#issuecomment-2195568292
|
||||
numpy<2.0.0
|
||||
@@ -0,0 +1,39 @@
|
||||
VIDEO_FOLDER="datasets/panda_70m/videos/data/"
|
||||
META_FILE_PATH="datasets/panda_70m/videos/meta_file_info.jsonl"
|
||||
SCENE_FOLDER="datasets/panda_70m/videos/meta_scene_info/"
|
||||
SCENE_SAVED_PATH="datasets/panda_70m/videos/meta_scene_info.jsonl"
|
||||
OUTPUT_FOLDER="datasets/panda_70m/videos_clips/data/"
|
||||
RESOLUTION_THRESHOLD=$((512*512))
|
||||
|
||||
# Set the duration range of video clips.
|
||||
export MIN_SECONDS=3
|
||||
export MAX_SECONDS=10
|
||||
|
||||
# Save all video names in a video folder as a meta file.
|
||||
python -m utils.get_meta_file \
|
||||
--video_folder $VIDEO_FOLDER \
|
||||
--saved_path $META_FILE_PATH
|
||||
|
||||
# Perform scene detection on the video dataset.
|
||||
# Adjust the n_jobs parameter based on the actual number of CPU cores in the machine.
|
||||
python cutscene_detect.py \
|
||||
--video_metadata_path $META_FILE_PATH \
|
||||
--video_folder $VIDEO_FOLDER \
|
||||
--saved_folder $SCENE_FOLDER \
|
||||
--n_jobs 32
|
||||
|
||||
# Gather all scene jsonl files to a single scene jsonl file.
|
||||
# Adjust the n_jobs parameter based on the actual I/O speed in the machine.
|
||||
python -m utils.gather_jsonl \
|
||||
--meta_folder $SCENE_FOLDER \
|
||||
--meta_file_path $SCENE_SAVED_PATH \
|
||||
--n_jobs 64
|
||||
|
||||
# Perform video splitting filtered by the RESOLUTION_THRESHOLD.
|
||||
# It consumes more CPU computing resources compared to the above operations.
|
||||
python video_splitting.py \
|
||||
--video_metadata_path $SCENE_SAVED_PATH \
|
||||
--video_folder $VIDEO_FOLDER \
|
||||
--output_folder $OUTPUT_FOLDER \
|
||||
--n_jobs 16 \
|
||||
--resolution_threshold $RESOLUTION_THRESHOLD
|
||||
@@ -0,0 +1,41 @@
|
||||
META_FILE_PATH="datasets/panda_70m/videos_clips/data/meta_file_info.jsonl"
|
||||
VIDEO_FOLDER="datasets/panda_70m/videos_clips/data/"
|
||||
VIDEO_QUALITY_SAVED_PATH="datasets/panda_70m/videos_clips/meta_quality_info_siglip.jsonl"
|
||||
MIN_ASETHETIC_SCORE_SIGLIP=4.0
|
||||
TEXT_SAVED_PATH="datasets/panda_70m/videos_clips/meta_text_info.jsonl"
|
||||
MIN_TEXT_SCORE=0.02
|
||||
MOTION_SAVED_PATH="datasets/panda_70m/videos_clips/meta_motion_info.jsonl"
|
||||
|
||||
python -m utils.get_meta_file \
|
||||
--video_folder $VIDEO_FOLDER \
|
||||
--saved_path $META_FILE_PATH
|
||||
|
||||
# Get the asethetic score (SigLIP) of all videos
|
||||
accelerate launch compute_video_quality.py \
|
||||
--video_metadata_path $META_FILE_PATH \
|
||||
--video_folder $VIDEO_FOLDER \
|
||||
--metrics "AestheticScoreSigLIP" \
|
||||
--frame_sample_method uniform \
|
||||
--num_sampled_frames 4 \
|
||||
--saved_freq 10 \
|
||||
--saved_path $VIDEO_QUALITY_SAVED_PATH \
|
||||
--batch_size 4
|
||||
|
||||
# Get the text score of all videos filtered by the video quality score.
|
||||
accelerate launch compute_text_score.py \
|
||||
--video_metadata_path $META_FILE_PATH \
|
||||
--video_folder $VIDEO_FOLDER \
|
||||
--saved_freq 10 \
|
||||
--saved_path $TEXT_SAVED_PATH \
|
||||
--asethetic_score_siglip_metadata_path $VIDEO_QUALITY_SAVED_PATH \
|
||||
--min_asethetic_score_siglip $MIN_ASETHETIC_SCORE_SIGLIP
|
||||
|
||||
# Get the motion score of all videos filtered by the video quality score and text score.
|
||||
python compute_motion_score.py \
|
||||
--video_metadata_path $META_FILE_PATH \
|
||||
--video_folder $VIDEO_FOLDER \
|
||||
--saved_freq 10 \
|
||||
--saved_path $MOTION_SAVED_PATH \
|
||||
--n_jobs 8 \
|
||||
--text_score_metadata_path $TEXT_SAVED_PATH \
|
||||
--min_text_score $MIN_TEXT_SCORE
|
||||
@@ -0,0 +1,52 @@
|
||||
META_FILE_PATH="datasets/panda_70m/videos_clips/data/meta_file_info.jsonl"
|
||||
VIDEO_FOLDER="datasets/panda_70m/videos_clips/data/"
|
||||
MOTION_SAVED_PATH="datasets/panda_70m/videos_clips/meta_motion_info.jsonl"
|
||||
MIN_MOTION_SCORE=2
|
||||
VIDEO_CAPTION_SAVED_PATH="datasets/panda_70m/meta_caption_info_vila_8b.jsonl"
|
||||
REWRITTEN_VIDEO_CAPTION_SAVED_PATH="datasets/panda_70m/meta_caption_info_vila_8b_rewritten.jsonl"
|
||||
VIDEOCLIPXL_SCORE_SAVED_PATH="datasets/panda_70m/meta_caption_info_vila_8b_rewritten_videoclipxl.jsonl"
|
||||
MIN_VIDEOCLIPXL_SCORE=0.20
|
||||
TRAIN_SAVED_PATH="datasets/panda_70m/train_panda_70m.json"
|
||||
# Manually download Efficient-Large-Model/Llama-3-VILA1.5-8b-AWQ to VILA_MODEL_PATH.
|
||||
# Manually download meta-llama/Meta-Llama-3-8B-Instruct to REWRITE_MODEL_PATH.
|
||||
|
||||
# Use VILA1.5-AWQ to perform recaptioning.
|
||||
accelerate launch vila_video_recaptioning.py \
|
||||
--video_metadata_path ${META_FILE_PATH} \
|
||||
--video_folder ${VIDEO_FOLDER} \
|
||||
--model_path ${VILA_MODEL_PATH} \
|
||||
--precision "W4A16" \
|
||||
--saved_path $VIDEO_CAPTION_SAVED_PATH \
|
||||
--saved_freq 1 \
|
||||
--motion_score_metadata_path $MOTION_SAVED_PATH \
|
||||
--min_motion_score $MIN_MOTION_SCORE
|
||||
|
||||
# Rewrite video captions (optional).
|
||||
python caption_rewrite.py \
|
||||
--video_metadata_path $VIDEO_CAPTION_SAVED_PATH \
|
||||
--batch_size 4096 \
|
||||
--model_name $REWRITE_MODEL_PATH \
|
||||
--prompt prompt/rewrite.txt \
|
||||
--prefix '"rewritten description": ' \
|
||||
--saved_path $REWRITTEN_VIDEO_CAPTION_SAVED_PATH \
|
||||
--saved_freq 1
|
||||
|
||||
# Compute caption-video alignment (optional).
|
||||
accelerate launch compute_video_quality.py \
|
||||
--video_metadata_path $REWRITTEN_VIDEO_CAPTION_SAVED_PATH \
|
||||
--caption_column caption \
|
||||
--video_folder $VIDEO_FOLDER \
|
||||
--frame_sample_method uniform \
|
||||
--num_sampled_frames 8 \
|
||||
--metrics VideoCLIPXLScore \
|
||||
--batch_size 4 \
|
||||
--saved_path $VIDEOCLIPXL_SCORE_SAVED_PATH \
|
||||
--saved_freq 10
|
||||
|
||||
# Get the final train file.
|
||||
python filter_meta_train.py \
|
||||
--caption_metadata_path $REWRITTEN_VIDEO_CAPTION_SAVED_PATH \
|
||||
--video_folder=$VIDEO_FOLDER \
|
||||
--videoclipxl_score_metadata_path $VIDEOCLIPXL_SCORE_SAVED_PATH \
|
||||
--min_videoclipxl_score $MIN_VIDEOCLIPXL_SCORE \
|
||||
--saved_path=$TRAIN_SAVED_PATH
|
||||
@@ -0,0 +1,162 @@
|
||||
import ast
|
||||
import os
|
||||
|
||||
import pandas as pd
|
||||
|
||||
from .logger import logger
|
||||
|
||||
|
||||
def filter(
|
||||
video_path_list,
|
||||
basic_metadata_path=None,
|
||||
min_resolution=0,
|
||||
min_duration=-1,
|
||||
max_duration=-1,
|
||||
asethetic_score_metadata_path=None,
|
||||
min_asethetic_score=4,
|
||||
asethetic_score_siglip_metadata_path=None,
|
||||
min_asethetic_score_siglip=4,
|
||||
text_score_metadata_path=None,
|
||||
min_text_score=0.02,
|
||||
motion_score_metadata_path=None,
|
||||
min_motion_score=2,
|
||||
videoclipxl_score_metadata_path=None,
|
||||
min_videoclipxl_score=0.20,
|
||||
video_path_column="video_path",
|
||||
):
|
||||
video_path_list = [os.path.basename(video_path) for video_path in video_path_list]
|
||||
|
||||
if basic_metadata_path is not None:
|
||||
if basic_metadata_path.endswith(".csv"):
|
||||
basic_df = pd.read_csv(basic_metadata_path)
|
||||
elif basic_metadata_path.endswith(".jsonl"):
|
||||
basic_df = pd.read_json(basic_metadata_path, lines=True)
|
||||
|
||||
basic_df["resolution"] = basic_df["frame_size"].apply(lambda x: x[0] * x[1])
|
||||
filtered_basic_df = basic_df[basic_df["resolution"] < min_resolution]
|
||||
filtered_video_path_list = filtered_basic_df[video_path_column].tolist()
|
||||
filtered_video_path_list = [os.path.basename(video_path) for video_path in filtered_video_path_list]
|
||||
|
||||
video_path_list = list(set(video_path_list).difference(set(filtered_video_path_list)))
|
||||
logger.info(
|
||||
f"Load {basic_metadata_path} ({len(basic_df)}) and filter {len(filtered_video_path_list)} videos "
|
||||
f"with resolution less than {min_resolution}."
|
||||
)
|
||||
|
||||
if min_duration != -1:
|
||||
filtered_basic_df = basic_df[basic_df["duration"] < min_duration]
|
||||
filtered_video_path_list = filtered_basic_df[video_path_column].tolist()
|
||||
filtered_video_path_list = [os.path.basename(video_path) for video_path in filtered_video_path_list]
|
||||
|
||||
video_path_list = list(set(video_path_list).difference(set(filtered_video_path_list)))
|
||||
logger.info(
|
||||
f"Load {basic_metadata_path} and filter {len(filtered_video_path_list)} videos "
|
||||
f"with duration less than {min_duration}."
|
||||
)
|
||||
|
||||
if max_duration != -1:
|
||||
filtered_basic_df = basic_df[basic_df["duration"] > max_duration]
|
||||
filtered_video_path_list = filtered_basic_df[video_path_column].tolist()
|
||||
filtered_video_path_list = [os.path.basename(video_path) for video_path in filtered_video_path_list]
|
||||
|
||||
video_path_list = list(set(video_path_list).difference(set(filtered_video_path_list)))
|
||||
logger.info(
|
||||
f"Load {basic_metadata_path} and filter {len(filtered_video_path_list)} videos "
|
||||
f"with duration greater than {max_duration}."
|
||||
)
|
||||
|
||||
if asethetic_score_metadata_path is not None:
|
||||
if asethetic_score_metadata_path.endswith(".csv"):
|
||||
asethetic_score_df = pd.read_csv(asethetic_score_metadata_path)
|
||||
elif asethetic_score_metadata_path.endswith(".jsonl"):
|
||||
asethetic_score_df = pd.read_json(asethetic_score_metadata_path, lines=True)
|
||||
|
||||
# In pandas, csv will save lists as strings, whereas jsonl will not.
|
||||
asethetic_score_df["aesthetic_score"] = asethetic_score_df["aesthetic_score"].apply(
|
||||
lambda x: ast.literal_eval(x) if isinstance(x, str) else x
|
||||
)
|
||||
asethetic_score_df["aesthetic_score_mean"] = asethetic_score_df["aesthetic_score"].apply(lambda x: sum(x) / len(x))
|
||||
filtered_asethetic_score_df = asethetic_score_df[asethetic_score_df["aesthetic_score_mean"] < min_asethetic_score]
|
||||
filtered_video_path_list = filtered_asethetic_score_df[video_path_column].tolist()
|
||||
filtered_video_path_list = [os.path.basename(video_path) for video_path in filtered_video_path_list]
|
||||
|
||||
video_path_list = list(set(video_path_list).difference(set(filtered_video_path_list)))
|
||||
logger.info(
|
||||
f"Load {asethetic_score_metadata_path} ({len(asethetic_score_df)}) and filter {len(filtered_video_path_list)} videos "
|
||||
f"with aesthetic score less than {min_asethetic_score}."
|
||||
)
|
||||
|
||||
if asethetic_score_siglip_metadata_path is not None:
|
||||
if asethetic_score_siglip_metadata_path.endswith(".csv"):
|
||||
asethetic_score_siglip_df = pd.read_csv(asethetic_score_siglip_metadata_path)
|
||||
elif asethetic_score_siglip_metadata_path.endswith(".jsonl"):
|
||||
asethetic_score_siglip_df = pd.read_json(asethetic_score_siglip_metadata_path, lines=True)
|
||||
|
||||
# In pandas, csv will save lists as strings, whereas jsonl will not.
|
||||
asethetic_score_siglip_df["aesthetic_score_siglip"] = asethetic_score_siglip_df["aesthetic_score_siglip"].apply(
|
||||
lambda x: ast.literal_eval(x) if isinstance(x, str) else x
|
||||
)
|
||||
asethetic_score_siglip_df["aesthetic_score_siglip_mean"] = asethetic_score_siglip_df["aesthetic_score_siglip"].apply(
|
||||
lambda x: sum(x) / len(x)
|
||||
)
|
||||
filtered_asethetic_score_siglip_df = asethetic_score_siglip_df[
|
||||
asethetic_score_siglip_df["aesthetic_score_siglip_mean"] < min_asethetic_score_siglip
|
||||
]
|
||||
filtered_video_path_list = filtered_asethetic_score_siglip_df[video_path_column].tolist()
|
||||
filtered_video_path_list = [os.path.basename(video_path) for video_path in filtered_video_path_list]
|
||||
|
||||
video_path_list = list(set(video_path_list).difference(set(filtered_video_path_list)))
|
||||
logger.info(
|
||||
f"Load {asethetic_score_siglip_metadata_path} ({len(asethetic_score_siglip_df)}) and filter {len(filtered_video_path_list)} videos "
|
||||
f"with aesthetic score (SigLIP) less than {min_asethetic_score_siglip}."
|
||||
)
|
||||
|
||||
if text_score_metadata_path is not None:
|
||||
if text_score_metadata_path.endswith(".csv"):
|
||||
text_score_df = pd.read_csv(text_score_metadata_path)
|
||||
elif text_score_metadata_path.endswith(".jsonl"):
|
||||
text_score_df = pd.read_json(text_score_metadata_path, lines=True)
|
||||
|
||||
filtered_text_score_df = text_score_df[text_score_df["text_score"] > min_text_score]
|
||||
filtered_video_path_list = filtered_text_score_df[video_path_column].tolist()
|
||||
filtered_video_path_list = [os.path.basename(video_path) for video_path in filtered_video_path_list]
|
||||
|
||||
video_path_list = list(set(video_path_list).difference(set(filtered_video_path_list)))
|
||||
logger.info(
|
||||
f"Load {text_score_metadata_path} ({len(text_score_df)}) and filter {len(filtered_video_path_list)} videos "
|
||||
f"with text score greater than {min_text_score}."
|
||||
)
|
||||
|
||||
if motion_score_metadata_path is not None:
|
||||
if motion_score_metadata_path.endswith(".csv"):
|
||||
motion_score_df = pd.read_csv(motion_score_metadata_path)
|
||||
elif motion_score_metadata_path.endswith(".jsonl"):
|
||||
motion_score_df = pd.read_json(motion_score_metadata_path, lines=True)
|
||||
|
||||
filtered_motion_score_df = motion_score_df[motion_score_df["motion_score"] < min_motion_score]
|
||||
filtered_video_path_list = filtered_motion_score_df[video_path_column].tolist()
|
||||
filtered_video_path_list = [os.path.basename(video_path) for video_path in filtered_video_path_list]
|
||||
|
||||
video_path_list = list(set(video_path_list).difference(set(filtered_video_path_list)))
|
||||
logger.info(
|
||||
f"Load {motion_score_metadata_path} ({len(motion_score_df)}) and filter {len(filtered_video_path_list)} videos "
|
||||
f"with motion score smaller than {min_motion_score}."
|
||||
)
|
||||
|
||||
if videoclipxl_score_metadata_path is not None:
|
||||
if videoclipxl_score_metadata_path.endswith(".csv"):
|
||||
videoclipxl_score_df = pd.read_csv(videoclipxl_score_metadata_path)
|
||||
elif videoclipxl_score_metadata_path.endswith(".jsonl"):
|
||||
videoclipxl_score_df = pd.read_json(videoclipxl_score_metadata_path, lines=True)
|
||||
|
||||
filtered_videoclipxl_score_df = videoclipxl_score_df[videoclipxl_score_df["videoclipxl_score"] < min_videoclipxl_score]
|
||||
filtered_video_path_list = filtered_videoclipxl_score_df[video_path_column].tolist()
|
||||
filtered_video_path_list = [os.path.basename(video_path) for video_path in filtered_video_path_list]
|
||||
|
||||
video_path_list = list(set(video_path_list).difference(set(filtered_video_path_list)))
|
||||
logger.info(
|
||||
f"Load {videoclipxl_score_metadata_path} ({len(videoclipxl_score_df)}) and "
|
||||
f"filter {len(filtered_video_path_list)} videos with mixclip score smaller than {min_videoclipxl_score}."
|
||||
)
|
||||
|
||||
return video_path_list
|
||||
@@ -0,0 +1,55 @@
|
||||
import argparse
|
||||
import os
|
||||
import glob
|
||||
import json
|
||||
from multiprocessing import Pool, Manager
|
||||
|
||||
import pandas as pd
|
||||
from natsort import index_natsorted
|
||||
|
||||
from .logger import logger
|
||||
|
||||
|
||||
def process_file(file_path, shared_list):
|
||||
with open(file_path, "r") as f:
|
||||
for line in f:
|
||||
data = json.loads(line)
|
||||
shared_list.append(data)
|
||||
|
||||
|
||||
def parse_args():
|
||||
parser = argparse.ArgumentParser(description="Gather all jsonl files in a folder (meta_folder) to a single jsonl file (meta_file_path).")
|
||||
parser.add_argument("--meta_folder", type=str, required=True)
|
||||
parser.add_argument("--meta_file_path", type=str, required=True)
|
||||
parser.add_argument("--video_path_column", type=str, default="video_path")
|
||||
parser.add_argument("--n_jobs", type=int, default=1)
|
||||
|
||||
args = parser.parse_args()
|
||||
return args
|
||||
|
||||
|
||||
def main():
|
||||
args = parse_args()
|
||||
|
||||
jsonl_files = glob.glob(os.path.join(args.meta_folder, "*.jsonl"))
|
||||
|
||||
with Manager() as manager:
|
||||
shared_list = manager.list()
|
||||
with Pool(processes=args.n_jobs) as pool:
|
||||
for file_path in jsonl_files:
|
||||
pool.apply_async(process_file, args=(file_path, shared_list))
|
||||
pool.close()
|
||||
pool.join()
|
||||
|
||||
with open(args.meta_file_path, "w") as f:
|
||||
for item in shared_list:
|
||||
f.write(json.dumps(item) + '\n')
|
||||
|
||||
df = pd.read_json(args.meta_file_path, lines=True)
|
||||
df = df.iloc[index_natsorted(df[args.video_path_column])].reset_index(drop=True)
|
||||
logger.info(f"Save the gathered single jsonl file to {args.meta_file_path}.")
|
||||
df.to_json(args.meta_file_path, orient="records", lines=True, force_ascii=False)
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
main()
|
||||
@@ -0,0 +1,74 @@
|
||||
import argparse
|
||||
from pathlib import Path
|
||||
|
||||
import pandas as pd
|
||||
from natsort import natsorted
|
||||
from tqdm import tqdm
|
||||
|
||||
from .logger import logger
|
||||
|
||||
|
||||
ALL_VIDEO_EXT = set(["mp4", "webm", "mkv", "avi", "flv", "mov"])
|
||||
ALL_IMGAE_EXT = set(["png", "webp", "jpg", "jpeg", "bmp", "gif"])
|
||||
|
||||
|
||||
def parse_args():
|
||||
parser = argparse.ArgumentParser(description="Compute scores of uniform sampled frames from videos.")
|
||||
parser.add_argument(
|
||||
"--image_path_column",
|
||||
type=str,
|
||||
default="image_path",
|
||||
help="The column contains the image path (an absolute path or a relative path w.r.t the image_folder).",
|
||||
)
|
||||
parser.add_argument("--image_folder", type=str, default=None, help="The video folder.")
|
||||
parser.add_argument(
|
||||
"--video_path_column",
|
||||
type=str,
|
||||
default="video_path",
|
||||
help="The column contains the video path (an absolute path or a relative path w.r.t the video_folder).",
|
||||
)
|
||||
parser.add_argument("--video_folder", type=str, default=None, help="The video folder.")
|
||||
parser.add_argument("--saved_path", type=str, required=True, help="The save path to the output results (csv/jsonl).")
|
||||
parser.add_argument("--recursive", action="store_true", help="Whether to search sub-folders recursively.")
|
||||
|
||||
args = parser.parse_args()
|
||||
return args
|
||||
|
||||
|
||||
def main():
|
||||
args = parse_args()
|
||||
|
||||
if args.video_folder is None and args.image_folder is None:
|
||||
raise ValueError("Either video_folder or image_folder should be specified in the arguments.")
|
||||
if args.video_folder is not None and args.image_folder is not None:
|
||||
raise ValueError("Both video_folder and image_folder can not be specified in the arguments at the same time.")
|
||||
|
||||
# Use the path name instead of the file name as video_path/image_path (unique ID).
|
||||
if args.video_folder is not None:
|
||||
video_path_list = []
|
||||
video_folder = Path(args.video_folder)
|
||||
for ext in tqdm(list(ALL_VIDEO_EXT)):
|
||||
if args.recursive:
|
||||
video_path_list += [str(file.relative_to(video_folder)) for file in video_folder.rglob(f"*.{ext}")]
|
||||
else:
|
||||
video_path_list += [str(file.relative_to(video_folder)) for file in video_folder.glob(f"*.{ext}")]
|
||||
video_path_list = natsorted(video_path_list)
|
||||
meta_file_df = pd.DataFrame({args.video_path_column: video_path_list})
|
||||
|
||||
if args.image_folder is not None:
|
||||
image_path_list = []
|
||||
image_folder = Path(args.image_folder)
|
||||
for ext in tqdm(list(ALL_IMGAE_EXT)):
|
||||
if args.recursive:
|
||||
image_path_list += [str(file.relative_to(image_folder)) for file in image_folder.rglob(f"*.{ext}")]
|
||||
else:
|
||||
image_path_list += [str(file.relative_to(image_folder)) for file in image_folder.glob(f"*.{ext}")]
|
||||
image_path_list = natsorted(image_path_list)
|
||||
meta_file_df = pd.DataFrame({args.image_path_column: image_path_list})
|
||||
|
||||
logger.info(f"{len(meta_file_df)} files in total. Save the result to {args.saved_path}.")
|
||||
meta_file_df.to_json(args.saved_path, orient="records", lines=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,248 @@
|
||||
import os
|
||||
from typing import Union
|
||||
|
||||
import clip
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from PIL import Image
|
||||
from torchvision.datasets.utils import download_url
|
||||
from transformers import AutoModel, AutoProcessor
|
||||
|
||||
from .siglip_v2_5 import convert_v2_5_from_siglip
|
||||
|
||||
# All metrics.
|
||||
__all__ = ["AestheticScore", "AestheticScoreSigLIP", "CLIPScore"]
|
||||
|
||||
_MODELS = {
|
||||
"CLIP_ViT-L/14": "https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/video_caption/clip/ViT-L-14.pt",
|
||||
"Aesthetics_V2": "https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/video_caption/clip/sac%2Blogos%2Bava1-l14-linearMSE.pth",
|
||||
"aesthetic_predictor_v2_5": "https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/video_caption/clip/aesthetic_predictor_v2_5.pth",
|
||||
}
|
||||
_MD5 = {
|
||||
"CLIP_ViT-L/14": "096db1af569b284eb76b3881534822d9",
|
||||
"Aesthetics_V2": "b1047fd767a00134b8fd6529bf19521a",
|
||||
"aesthetic_predictor_v2_5": "c46eb8c29f714c9231dc630b8226842a",
|
||||
}
|
||||
|
||||
|
||||
def get_list_depth(lst):
|
||||
if isinstance(lst, list):
|
||||
return 1 + max(get_list_depth(item) for item in lst)
|
||||
else:
|
||||
return 0
|
||||
|
||||
|
||||
def reshape_images(images: Union[list[list[Image.Image]], list[Image.Image]]):
|
||||
# Check the input sanity.
|
||||
depth = get_list_depth(images)
|
||||
if depth == 1: # batch image input
|
||||
if not isinstance(images[0], Image.Image):
|
||||
raise ValueError("The item in 1D images should be Image.Image.")
|
||||
num_sampled_frames = None
|
||||
elif depth == 2: # batch video input
|
||||
if not isinstance(images[0][0], Image.Image):
|
||||
raise ValueError("The item in 2D images (videos) should be Image.Image.")
|
||||
num_sampled_frames = len(images[0])
|
||||
if not all(len(video_frames) == num_sampled_frames for video_frames in images):
|
||||
raise ValueError("All item in 2D images should be with the same length.")
|
||||
# [batch_size, num_sampled_frames, H, W, C] => [batch_size * num_sampled_frames, H, W, C].
|
||||
reshaped_images = []
|
||||
for video_frames in images:
|
||||
reshaped_images.extend([frame for frame in video_frames])
|
||||
images = reshaped_images
|
||||
else:
|
||||
raise ValueError("The input images should be in 1/2D list.")
|
||||
|
||||
return images, num_sampled_frames
|
||||
|
||||
|
||||
def reshape_scores(scores: list[float], num_sampled_frames: int) -> list[float]:
|
||||
if isinstance(scores, list):
|
||||
if num_sampled_frames is not None: # Batch video input
|
||||
batch_size = len(scores) // num_sampled_frames
|
||||
scores = [
|
||||
scores[i * num_sampled_frames:(i + 1) * num_sampled_frames]
|
||||
for i in range(batch_size)
|
||||
]
|
||||
return scores
|
||||
else:
|
||||
return [scores]
|
||||
|
||||
|
||||
# if you changed the MLP architecture during training, change it also here:
|
||||
class _MLP(nn.Module):
|
||||
def __init__(self, input_size):
|
||||
super().__init__()
|
||||
self.input_size = input_size
|
||||
self.layers = nn.Sequential(
|
||||
nn.Linear(self.input_size, 1024),
|
||||
# nn.ReLU(),
|
||||
nn.Dropout(0.2),
|
||||
nn.Linear(1024, 128),
|
||||
# nn.ReLU(),
|
||||
nn.Dropout(0.2),
|
||||
nn.Linear(128, 64),
|
||||
# nn.ReLU(),
|
||||
nn.Dropout(0.1),
|
||||
nn.Linear(64, 16),
|
||||
# nn.ReLU(),
|
||||
nn.Linear(16, 1),
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
return self.layers(x)
|
||||
|
||||
|
||||
class AestheticScore:
|
||||
"""Compute LAION Aesthetics Score V2 based on openai/clip. Note that the default
|
||||
inference dtype with GPUs is fp16 in openai/clip.
|
||||
|
||||
Ref:
|
||||
1. https://github.com/christophschuhmann/improved-aesthetic-predictor/blob/main/simple_inference.py.
|
||||
2. https://github.com/openai/CLIP/issues/30.
|
||||
"""
|
||||
|
||||
def __init__(self, root: str = "~/.cache/clip", device: str = "cpu"):
|
||||
# The CLIP model is loaded in the evaluation mode.
|
||||
self.root = os.path.expanduser(root)
|
||||
if not os.path.exists(self.root):
|
||||
os.makedirs(self.root)
|
||||
filename = "ViT-L-14.pt"
|
||||
download_url(_MODELS["CLIP_ViT-L/14"], self.root, filename=filename, md5=_MD5["CLIP_ViT-L/14"])
|
||||
self.clip_model, self.preprocess = clip.load(os.path.join(self.root, filename), device=device)
|
||||
self.device = device
|
||||
self._load_mlp()
|
||||
|
||||
def _load_mlp(self):
|
||||
filename = "sac+logos+ava1-l14-linearMSE.pth"
|
||||
download_url(_MODELS["Aesthetics_V2"], self.root, filename=filename, md5=_MD5["Aesthetics_V2"])
|
||||
state_dict = torch.load(os.path.join(self.root, filename))
|
||||
self.mlp = _MLP(768)
|
||||
self.mlp.load_state_dict(state_dict)
|
||||
self.mlp.to(self.device)
|
||||
self.mlp.eval()
|
||||
|
||||
def __call__(self, images: Union[list[list[Image.Image]], list[Image.Image]], texts=None) -> list[float]:
|
||||
images, num_sampled_frames = reshape_images(images)
|
||||
|
||||
with torch.no_grad():
|
||||
images = torch.stack([self.preprocess(image) for image in images]).to(self.device)
|
||||
image_embs = F.normalize(self.clip_model.encode_image(images))
|
||||
scores = self.mlp(image_embs.float()) # torch.float16 -> torch.float32, [N, 1]
|
||||
|
||||
scores = scores.squeeze().tolist() # scalar or list
|
||||
return reshape_scores(scores, num_sampled_frames)
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return "aesthetic_score"
|
||||
|
||||
|
||||
class AestheticScoreSigLIP:
|
||||
"""Compute Aesthetics Score V2.5 based on google/siglip-so400m-patch14-384.
|
||||
|
||||
Ref:
|
||||
1. https://github.com/discus0434/aesthetic-predictor-v2-5.
|
||||
2. https://github.com/discus0434/aesthetic-predictor-v2-5/issues/2.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
root: str = "~/.cache/clip",
|
||||
device: str = "cpu",
|
||||
torch_dtype=torch.float16
|
||||
):
|
||||
self.root = os.path.expanduser(root)
|
||||
if not os.path.exists(self.root):
|
||||
os.makedirs(self.root)
|
||||
filename = "aesthetic_predictor_v2_5.pth"
|
||||
download_url(_MODELS["aesthetic_predictor_v2_5"], self.root, filename=filename, md5=_MD5["aesthetic_predictor_v2_5"])
|
||||
self.model, self.preprocessor = convert_v2_5_from_siglip(
|
||||
predictor_name_or_path=os.path.join(self.root, filename),
|
||||
low_cpu_mem_usage=True,
|
||||
trust_remote_code=True,
|
||||
)
|
||||
self.model = self.model.to(device=device, dtype=torch_dtype)
|
||||
self.device = device
|
||||
self.torch_dtype = torch_dtype
|
||||
|
||||
def __call__(self, images: Union[list[list[Image.Image]], list[Image.Image]], texts=None) -> list[float]:
|
||||
images, num_sampled_frames = reshape_images(images)
|
||||
|
||||
pixel_values = self.preprocessor(images, return_tensors="pt").pixel_values
|
||||
pixel_values = pixel_values.to(self.device, self.torch_dtype)
|
||||
with torch.no_grad():
|
||||
scores = self.model(pixel_values).logits.squeeze().float().cpu().numpy()
|
||||
|
||||
scores = scores.squeeze().tolist() # scalar or list
|
||||
return reshape_scores(scores, num_sampled_frames)
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return "aesthetic_score_siglip"
|
||||
|
||||
|
||||
class CLIPScore:
|
||||
"""Compute CLIP scores for image-text pairs based on huggingface/transformers."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model_name_or_path: str = "openai/clip-vit-large-patch14",
|
||||
torch_dtype=torch.float16,
|
||||
device: str = "cpu",
|
||||
):
|
||||
self.model = AutoModel.from_pretrained(model_name_or_path, torch_dtype=torch_dtype).eval().to(device)
|
||||
self.processor = AutoProcessor.from_pretrained(model_name_or_path)
|
||||
self.torch_dtype = torch_dtype
|
||||
self.device = device
|
||||
|
||||
def __call__(self, images: Union[list[list[Image.Image]], list[Image.Image]], texts: list[str]) -> list[float]:
|
||||
assert len(images) == len(texts)
|
||||
images, num_sampled_frames = reshape_images(images)
|
||||
# Expand texts in the batch video input case.
|
||||
if num_sampled_frames is not None:
|
||||
texts = [[text] * num_sampled_frames for text in texts]
|
||||
texts = [item for sublist in texts for item in sublist]
|
||||
|
||||
image_inputs = self.processor(images=images, return_tensors="pt") # {"pixel_values": }
|
||||
if self.torch_dtype == torch.float16:
|
||||
image_inputs["pixel_values"] = image_inputs["pixel_values"].half()
|
||||
text_inputs = self.processor(text=texts, return_tensors="pt", padding=True, truncation=True) # {"inputs_id": }
|
||||
image_inputs, text_inputs = image_inputs.to(self.device), text_inputs.to(self.device)
|
||||
with torch.no_grad():
|
||||
image_embs = F.normalize(self.model.get_image_features(**image_inputs))
|
||||
text_embs = F.normalize(self.model.get_text_features(**text_inputs))
|
||||
scores = text_embs @ image_embs.T # [N, N]
|
||||
|
||||
scores = scores.squeeze().tolist() # scalar or list
|
||||
return reshape_scores(scores, num_sampled_frames)
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return "clip_score"
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
from torch.utils.data import DataLoader
|
||||
from tqdm import tqdm
|
||||
from .video_dataset import VideoDataset, collate_fn
|
||||
|
||||
aesthetic_score = AestheticScore(device="cuda")
|
||||
aesthetic_score_siglip = AestheticScoreSigLIP(device="cuda")
|
||||
# clip_score = CLIPScore(device="cuda")
|
||||
|
||||
paths = ["your_image_path"] * 3
|
||||
# texts = ["a joker", "a woman", "a man"]
|
||||
images = [Image.open(p).convert("RGB") for p in paths]
|
||||
|
||||
print(aesthetic_score(images))
|
||||
# print(clip_score(images, texts))
|
||||
|
||||
test_dataset = VideoDataset(
|
||||
dataset_inputs={"video_path": ["your_video_path"] * 3},
|
||||
sample_method="mid",
|
||||
num_sampled_frames=2
|
||||
)
|
||||
test_loader = DataLoader(test_dataset, batch_size=1, num_workers=1, collate_fn=collate_fn)
|
||||
|
||||
for idx, batch in enumerate(tqdm(test_loader)):
|
||||
batch_frame = batch["sampled_frame"]
|
||||
print(aesthetic_score_siglip(batch_frame))
|
||||
@@ -0,0 +1,36 @@
|
||||
# Borrowed from sd-webui-controlnet/scripts/logging.py
|
||||
import copy
|
||||
import logging
|
||||
import sys
|
||||
|
||||
|
||||
class ColoredFormatter(logging.Formatter):
|
||||
COLORS = {
|
||||
"DEBUG": "\033[0;36m", # CYAN
|
||||
"INFO": "\033[0;32m", # GREEN
|
||||
"WARNING": "\033[0;33m", # YELLOW
|
||||
"ERROR": "\033[0;31m", # RED
|
||||
"CRITICAL": "\033[0;37;41m", # WHITE ON RED
|
||||
"RESET": "\033[0m", # RESET COLOR
|
||||
}
|
||||
|
||||
def format(self, record):
|
||||
colored_record = copy.copy(record)
|
||||
levelname = colored_record.levelname
|
||||
seq = self.COLORS.get(levelname, self.COLORS["RESET"])
|
||||
colored_record.levelname = f"{seq}{levelname}{self.COLORS['RESET']}"
|
||||
return super().format(colored_record)
|
||||
|
||||
|
||||
# Create a new logger
|
||||
logger = logging.getLogger("VideoCaption")
|
||||
logger.propagate = False
|
||||
|
||||
# Add handler if we don't have one.
|
||||
if not logger.handlers:
|
||||
handler = logging.StreamHandler(sys.stdout)
|
||||
handler.setFormatter(ColoredFormatter("%(asctime)s - %(name)s - %(levelname)s - %(message)s"))
|
||||
logger.addHandler(handler)
|
||||
|
||||
# Configure logger
|
||||
logger.setLevel("INFO")
|
||||
@@ -0,0 +1,19 @@
|
||||
# Long-CLIP
|
||||
Codes in this directory are borrowed from https://github.com/beichenzbc/Long-CLIP/tree/4e6f5da/model.
|
||||
|
||||
We only modify the following code in [model_longclip.py](model_longclip.py) from
|
||||
```python
|
||||
@property
|
||||
def dtype(self):
|
||||
return self.visual.conv1.weight.dtype
|
||||
```
|
||||
to
|
||||
```python
|
||||
@property
|
||||
def dtype(self):
|
||||
# Fix: the VideoCLIP-XL inference.
|
||||
if hasattr(self, "visual"):
|
||||
return self.visual.conv1.weight.dtype
|
||||
else:
|
||||
return self.token_embedding.weight.dtype
|
||||
```
|
||||
@@ -0,0 +1 @@
|
||||
from .longclip import *
|
||||
Binary file not shown.
@@ -0,0 +1,353 @@
|
||||
import hashlib
|
||||
import os
|
||||
import urllib
|
||||
import warnings
|
||||
from typing import Any, Union, List
|
||||
from pkg_resources import packaging
|
||||
from torch import nn
|
||||
import torch
|
||||
from PIL import Image
|
||||
from torchvision.transforms import Compose, Resize, CenterCrop, ToTensor, Normalize
|
||||
from tqdm import tqdm
|
||||
|
||||
from .model_longclip import build_model
|
||||
from .simple_tokenizer import SimpleTokenizer as _Tokenizer
|
||||
|
||||
try:
|
||||
from torchvision.transforms import InterpolationMode
|
||||
BICUBIC = InterpolationMode.BICUBIC
|
||||
except ImportError:
|
||||
BICUBIC = Image.BICUBIC
|
||||
|
||||
|
||||
if packaging.version.parse(torch.__version__) < packaging.version.parse("1.7.1"):
|
||||
warnings.warn("PyTorch version 1.7.1 or higher is recommended")
|
||||
|
||||
|
||||
__all__ = ["load", "tokenize"]
|
||||
_tokenizer = _Tokenizer()
|
||||
|
||||
|
||||
def _convert_image_to_rgb(image):
|
||||
return image.convert("RGB")
|
||||
|
||||
|
||||
def _transform(n_px):
|
||||
return Compose([
|
||||
Resize(n_px, interpolation=BICUBIC),
|
||||
CenterCrop(n_px),
|
||||
_convert_image_to_rgb,
|
||||
ToTensor(),
|
||||
Normalize((0.48145466, 0.4578275, 0.40821073), (0.26862954, 0.26130258, 0.27577711)),
|
||||
])
|
||||
|
||||
|
||||
|
||||
def load(name: str, device: Union[str, torch.device] = "cuda" if torch.cuda.is_available() else "cpu", download_root: str = None):
|
||||
"""Load a long CLIP model
|
||||
|
||||
Parameters
|
||||
----------
|
||||
name : str
|
||||
A model name listed by `clip.available_models()`, or the path to a model checkpoint containing the state_dict
|
||||
|
||||
device : Union[str, torch.device]
|
||||
The device to put the loaded model
|
||||
|
||||
Returns
|
||||
-------
|
||||
model : torch.nn.Module
|
||||
The CLIP model
|
||||
|
||||
preprocess : Callable[[PIL.Image], torch.Tensor]
|
||||
A torchvision transform that converts a PIL image into a tensor that the returned model can take as its input
|
||||
"""
|
||||
|
||||
model_path = name
|
||||
|
||||
state_dict = torch.load(model_path, map_location="cpu")
|
||||
|
||||
model = build_model(state_dict or model.state_dict(), load_from_clip = False).to(device)
|
||||
|
||||
if str(device) == "cpu":
|
||||
model.float()
|
||||
|
||||
return model, _transform(model.visual.input_resolution)
|
||||
|
||||
|
||||
|
||||
def _node_get(node: torch._C.Node, key: str):
|
||||
"""Gets attributes of a node which is polymorphic over return type.
|
||||
|
||||
From https://github.com/pytorch/pytorch/pull/82628
|
||||
"""
|
||||
sel = node.kindOf(key)
|
||||
return getattr(node, sel)(key)
|
||||
|
||||
def patch_device(module):
|
||||
try:
|
||||
graphs = [module.graph] if hasattr(module, "graph") else []
|
||||
except RuntimeError:
|
||||
graphs = []
|
||||
|
||||
if hasattr(module, "forward1"):
|
||||
graphs.append(module.forward1.graph)
|
||||
|
||||
for graph in graphs:
|
||||
for node in graph.findAllNodes("prim::Constant"):
|
||||
if "value" in node.attributeNames() and str(_node_get(node, "value")).startswith("cuda"):
|
||||
node.copyAttributes(device_node)
|
||||
|
||||
model.apply(patch_device)
|
||||
patch_device(model.encode_image)
|
||||
patch_device(model.encode_text)
|
||||
|
||||
# patch dtype to float32 on CPU
|
||||
if str(device) == "cpu":
|
||||
float_holder = torch.jit.trace(lambda: torch.ones([]).float(), example_inputs=[])
|
||||
float_input = list(float_holder.graph.findNode("aten::to").inputs())[1]
|
||||
float_node = float_input.node()
|
||||
|
||||
def patch_float(module):
|
||||
try:
|
||||
graphs = [module.graph] if hasattr(module, "graph") else []
|
||||
except RuntimeError:
|
||||
graphs = []
|
||||
|
||||
if hasattr(module, "forward1"):
|
||||
graphs.append(module.forward1.graph)
|
||||
|
||||
for graph in graphs:
|
||||
for node in graph.findAllNodes("aten::to"):
|
||||
inputs = list(node.inputs())
|
||||
for i in [1, 2]: # dtype can be the second or third argument to aten::to()
|
||||
if _node_get(inputs[i].node(), "value") == 5:
|
||||
inputs[i].node().copyAttributes(float_node)
|
||||
|
||||
model.apply(patch_float)
|
||||
patch_float(model.encode_image)
|
||||
patch_float(model.encode_text)
|
||||
|
||||
model.float()
|
||||
|
||||
return model, _transform(model.input_resolution.item())
|
||||
|
||||
|
||||
def load_from_clip(name: str, device: Union[str, torch.device] = "cuda" if torch.cuda.is_available() else "cpu", jit: bool = False, download_root: str = None):
|
||||
"""Load from CLIP model for fine-tuning
|
||||
|
||||
Parameters
|
||||
----------
|
||||
name : str
|
||||
A model name listed by `clip.available_models()`, or the path to a model checkpoint containing the state_dict
|
||||
|
||||
device : Union[str, torch.device]
|
||||
The device to put the loaded model
|
||||
|
||||
jit : bool
|
||||
Whether to load the optimized JIT model or more hackable non-JIT model (default).
|
||||
|
||||
download_root: str
|
||||
path to download the model files; by default, it uses "~/.cache/clip"
|
||||
|
||||
Returns
|
||||
-------
|
||||
model : torch.nn.Module
|
||||
The CLIP model
|
||||
|
||||
preprocess : Callable[[PIL.Image], torch.Tensor]
|
||||
A torchvision transform that converts a PIL image into a tensor that the returned model can take as its input
|
||||
"""
|
||||
|
||||
_MODELS = {
|
||||
"RN50": "https://openaipublic.azureedge.net/clip/models/afeb0e10f9e5a86da6080e35cf09123aca3b358a0c3e3b6c78a7b63bc04b6762/RN50.pt",
|
||||
"RN101": "https://openaipublic.azureedge.net/clip/models/8fa8567bab74a42d41c5915025a8e4538c3bdbe8804a470a72f30b0d94fab599/RN101.pt",
|
||||
"RN50x4": "https://openaipublic.azureedge.net/clip/models/7e526bd135e493cef0776de27d5f42653e6b4c8bf9e0f653bb11773263205fdd/RN50x4.pt",
|
||||
"RN50x16": "https://openaipublic.azureedge.net/clip/models/52378b407f34354e150460fe41077663dd5b39c54cd0bfd2b27167a4a06ec9aa/RN50x16.pt",
|
||||
"RN50x64": "https://openaipublic.azureedge.net/clip/models/be1cfb55d75a9666199fb2206c106743da0f6468c9d327f3e0d0a543a9919d9c/RN50x64.pt",
|
||||
"ViT-B/32": "https://openaipublic.azureedge.net/clip/models/40d365715913c9da98579312b702a82c18be219cc2a73407c4526f58eba950af/ViT-B-32.pt",
|
||||
"ViT-B/16": "https://openaipublic.azureedge.net/clip/models/5806e77cd80f8b59890b7e101eabd078d9fb84e6937f9e85e4ecb61988df416f/ViT-B-16.pt",
|
||||
"ViT-L/14": "https://openaipublic.azureedge.net/clip/models/b8cca3fd41ae0c99ba7e8951adf17d267cdb84cd88be6f7c2e0eca1737a03836/ViT-L-14.pt",
|
||||
"ViT-L/14@336px": "https://openaipublic.azureedge.net/clip/models/3035c92b350959924f9f00213499208652fc7ea050643e8b385c2dac08641f02/ViT-L-14-336px.pt",
|
||||
}
|
||||
|
||||
def available_models() -> List[str]:
|
||||
"""Returns the names of available CLIP models"""
|
||||
return list(_MODELS.keys())
|
||||
|
||||
def _download(url: str, root: str):
|
||||
os.makedirs(root, exist_ok=True)
|
||||
filename = os.path.basename(url)
|
||||
|
||||
expected_sha256 = url.split("/")[-2]
|
||||
download_target = os.path.join(root, filename)
|
||||
|
||||
if os.path.exists(download_target) and not os.path.isfile(download_target):
|
||||
raise RuntimeError(f"{download_target} exists and is not a regular file")
|
||||
|
||||
if os.path.isfile(download_target):
|
||||
if hashlib.sha256(open(download_target, "rb").read()).hexdigest() == expected_sha256:
|
||||
return download_target
|
||||
else:
|
||||
warnings.warn(f"{download_target} exists, but the SHA256 checksum does not match; re-downloading the file")
|
||||
|
||||
with urllib.request.urlopen(url) as source, open(download_target, "wb") as output:
|
||||
with tqdm(total=int(source.info().get("Content-Length")), ncols=80, unit='iB', unit_scale=True, unit_divisor=1024) as loop:
|
||||
while True:
|
||||
buffer = source.read(8192)
|
||||
if not buffer:
|
||||
break
|
||||
|
||||
output.write(buffer)
|
||||
loop.update(len(buffer))
|
||||
|
||||
if hashlib.sha256(open(download_target, "rb").read()).hexdigest() != expected_sha256:
|
||||
raise RuntimeError("Model has been downloaded but the SHA256 checksum does not not match")
|
||||
|
||||
return download_target
|
||||
|
||||
if name in _MODELS:
|
||||
model_path = _download(_MODELS[name], download_root or os.path.expanduser("~/.cache/clip"))
|
||||
elif os.path.isfile(name):
|
||||
model_path = name
|
||||
else:
|
||||
raise RuntimeError(f"Model {name} not found; available models = {available_models()}")
|
||||
|
||||
with open(model_path, 'rb') as opened_file:
|
||||
try:
|
||||
# loading JIT archive
|
||||
model = torch.jit.load(opened_file, map_location=device if jit else "cpu").eval()
|
||||
state_dict = None
|
||||
except RuntimeError:
|
||||
# loading saved state dict
|
||||
if jit:
|
||||
warnings.warn(f"File {model_path} is not a JIT archive. Loading as a state dict instead")
|
||||
jit = False
|
||||
state_dict = torch.load(opened_file, map_location="cpu")
|
||||
|
||||
model = build_model(state_dict or model.state_dict(), load_from_clip = True).to(device)
|
||||
|
||||
positional_embedding_pre = model.positional_embedding.type(model.dtype)
|
||||
|
||||
length, dim = positional_embedding_pre.shape
|
||||
keep_len = 20
|
||||
posisitonal_embedding_new = torch.zeros([4*length-3*keep_len, dim], dtype=model.dtype)
|
||||
for i in range(keep_len):
|
||||
posisitonal_embedding_new[i] = positional_embedding_pre[i]
|
||||
for i in range(length-1-keep_len):
|
||||
posisitonal_embedding_new[4*i + keep_len] = positional_embedding_pre[i + keep_len]
|
||||
posisitonal_embedding_new[4*i + 1 + keep_len] = 3*positional_embedding_pre[i + keep_len]/4 + 1*positional_embedding_pre[i+1+keep_len]/4
|
||||
posisitonal_embedding_new[4*i + 2+keep_len] = 2*positional_embedding_pre[i+keep_len]/4 + 2*positional_embedding_pre[i+1+keep_len]/4
|
||||
posisitonal_embedding_new[4*i + 3+keep_len] = 1*positional_embedding_pre[i+keep_len]/4 + 3*positional_embedding_pre[i+1+keep_len]/4
|
||||
|
||||
posisitonal_embedding_new[4*length -3*keep_len - 4] = positional_embedding_pre[length-1] + 0*(positional_embedding_pre[length-1] - positional_embedding_pre[length-2])/4
|
||||
posisitonal_embedding_new[4*length -3*keep_len - 3] = positional_embedding_pre[length-1] + 1*(positional_embedding_pre[length-1] - positional_embedding_pre[length-2])/4
|
||||
posisitonal_embedding_new[4*length -3*keep_len - 2] = positional_embedding_pre[length-1] + 2*(positional_embedding_pre[length-1] - positional_embedding_pre[length-2])/4
|
||||
posisitonal_embedding_new[4*length -3*keep_len - 1] = positional_embedding_pre[length-1] + 3*(positional_embedding_pre[length-1] - positional_embedding_pre[length-2])/4
|
||||
|
||||
positional_embedding_res = posisitonal_embedding_new.clone()
|
||||
|
||||
model.positional_embedding = nn.Parameter(posisitonal_embedding_new, requires_grad=False)
|
||||
model.positional_embedding_res = nn.Parameter(positional_embedding_res, requires_grad=True)
|
||||
|
||||
if str(device) == "cpu":
|
||||
model.float()
|
||||
return model, _transform(model.visual.input_resolution)
|
||||
|
||||
def _node_get(node: torch._C.Node, key: str):
|
||||
"""Gets attributes of a node which is polymorphic over return type.
|
||||
|
||||
From https://github.com/pytorch/pytorch/pull/82628
|
||||
"""
|
||||
sel = node.kindOf(key)
|
||||
return getattr(node, sel)(key)
|
||||
|
||||
def patch_device(module):
|
||||
try:
|
||||
graphs = [module.graph] if hasattr(module, "graph") else []
|
||||
except RuntimeError:
|
||||
graphs = []
|
||||
|
||||
if hasattr(module, "forward1"):
|
||||
graphs.append(module.forward1.graph)
|
||||
|
||||
for graph in graphs:
|
||||
for node in graph.findAllNodes("prim::Constant"):
|
||||
if "value" in node.attributeNames() and str(_node_get(node, "value")).startswith("cuda"):
|
||||
node.copyAttributes(device_node)
|
||||
|
||||
model.apply(patch_device)
|
||||
patch_device(model.encode_image)
|
||||
patch_device(model.encode_text)
|
||||
|
||||
# patch dtype to float32 on CPU
|
||||
if str(device) == "cpu":
|
||||
float_holder = torch.jit.trace(lambda: torch.ones([]).float(), example_inputs=[])
|
||||
float_input = list(float_holder.graph.findNode("aten::to").inputs())[1]
|
||||
float_node = float_input.node()
|
||||
|
||||
def patch_float(module):
|
||||
try:
|
||||
graphs = [module.graph] if hasattr(module, "graph") else []
|
||||
except RuntimeError:
|
||||
graphs = []
|
||||
|
||||
if hasattr(module, "forward1"):
|
||||
graphs.append(module.forward1.graph)
|
||||
|
||||
for graph in graphs:
|
||||
for node in graph.findAllNodes("aten::to"):
|
||||
inputs = list(node.inputs())
|
||||
for i in [1, 2]: # dtype can be the second or third argument to aten::to()
|
||||
if _node_get(inputs[i].node(), "value") == 5:
|
||||
inputs[i].node().copyAttributes(float_node)
|
||||
|
||||
model.apply(patch_float)
|
||||
patch_float(model.encode_image)
|
||||
patch_float(model.encode_text)
|
||||
|
||||
model.float()
|
||||
|
||||
return model, _transform(model.input_resolution.item())
|
||||
|
||||
def tokenize(texts: Union[str, List[str]], context_length: int = 77*4-60, truncate: bool = False) -> Union[torch.IntTensor, torch.LongTensor]:
|
||||
"""
|
||||
Returns the tokenized representation of given input string(s)
|
||||
|
||||
Parameters
|
||||
----------
|
||||
texts : Union[str, List[str]]
|
||||
An input string or a list of input strings to tokenize
|
||||
|
||||
context_length : int
|
||||
The context length to use; all CLIP models use 77 as the context length
|
||||
|
||||
truncate: bool
|
||||
Whether to truncate the text in case its encoding is longer than the context length
|
||||
|
||||
Returns
|
||||
-------
|
||||
A two-dimensional tensor containing the resulting tokens, shape = [number of input strings, context_length].
|
||||
We return LongTensor when torch version is <1.8.0, since older index_select requires indices to be long.
|
||||
"""
|
||||
if isinstance(texts, str):
|
||||
texts = [texts]
|
||||
|
||||
sot_token = _tokenizer.encoder["<|startoftext|>"]
|
||||
eot_token = _tokenizer.encoder["<|endoftext|>"]
|
||||
all_tokens = [[sot_token] + _tokenizer.encode(text) + [eot_token] for text in texts]
|
||||
if packaging.version.parse(torch.__version__) < packaging.version.parse("1.8.0"):
|
||||
result = torch.zeros(len(all_tokens), context_length, dtype=torch.long)
|
||||
else:
|
||||
result = torch.zeros(len(all_tokens), context_length, dtype=torch.int)
|
||||
|
||||
for i, tokens in enumerate(all_tokens):
|
||||
if len(tokens) > context_length:
|
||||
if truncate:
|
||||
tokens = tokens[:context_length]
|
||||
tokens[-1] = eot_token
|
||||
else:
|
||||
raise RuntimeError(f"Input {texts[i]} is too long for context length {context_length}")
|
||||
result[i, :len(tokens)] = torch.tensor(tokens)
|
||||
|
||||
return result
|
||||
@@ -0,0 +1,471 @@
|
||||
from collections import OrderedDict
|
||||
from typing import Tuple, Union
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from torch import nn
|
||||
|
||||
|
||||
class Bottleneck(nn.Module):
|
||||
expansion = 4
|
||||
|
||||
def __init__(self, inplanes, planes, stride=1):
|
||||
super().__init__()
|
||||
|
||||
# all conv layers have stride 1. an avgpool is performed after the second convolution when stride > 1
|
||||
self.conv1 = nn.Conv2d(inplanes, planes, 1, bias=False)
|
||||
self.bn1 = nn.BatchNorm2d(planes)
|
||||
self.relu1 = nn.ReLU(inplace=True)
|
||||
|
||||
self.conv2 = nn.Conv2d(planes, planes, 3, padding=1, bias=False)
|
||||
self.bn2 = nn.BatchNorm2d(planes)
|
||||
self.relu2 = nn.ReLU(inplace=True)
|
||||
|
||||
self.avgpool = nn.AvgPool2d(stride) if stride > 1 else nn.Identity()
|
||||
|
||||
self.conv3 = nn.Conv2d(planes, planes * self.expansion, 1, bias=False)
|
||||
self.bn3 = nn.BatchNorm2d(planes * self.expansion)
|
||||
self.relu3 = nn.ReLU(inplace=True)
|
||||
|
||||
self.downsample = None
|
||||
self.stride = stride
|
||||
|
||||
if stride > 1 or inplanes != planes * Bottleneck.expansion:
|
||||
# downsampling layer is prepended with an avgpool, and the subsequent convolution has stride 1
|
||||
self.downsample = nn.Sequential(OrderedDict([
|
||||
("-1", nn.AvgPool2d(stride)),
|
||||
("0", nn.Conv2d(inplanes, planes * self.expansion, 1, stride=1, bias=False)),
|
||||
("1", nn.BatchNorm2d(planes * self.expansion))
|
||||
]))
|
||||
|
||||
def forward(self, x: torch.Tensor):
|
||||
identity = x
|
||||
|
||||
out = self.relu1(self.bn1(self.conv1(x)))
|
||||
out = self.relu2(self.bn2(self.conv2(out)))
|
||||
out = self.avgpool(out)
|
||||
out = self.bn3(self.conv3(out))
|
||||
|
||||
if self.downsample is not None:
|
||||
identity = self.downsample(x)
|
||||
|
||||
out += identity
|
||||
out = self.relu3(out)
|
||||
return out
|
||||
|
||||
|
||||
class AttentionPool2d(nn.Module):
|
||||
def __init__(self, spacial_dim: int, embed_dim: int, num_heads: int, output_dim: int = None):
|
||||
super().__init__()
|
||||
self.positional_embedding = nn.Parameter(torch.randn(spacial_dim ** 2 + 1, embed_dim) / embed_dim ** 0.5)
|
||||
self.k_proj = nn.Linear(embed_dim, embed_dim)
|
||||
self.q_proj = nn.Linear(embed_dim, embed_dim)
|
||||
self.v_proj = nn.Linear(embed_dim, embed_dim)
|
||||
self.c_proj = nn.Linear(embed_dim, output_dim or embed_dim)
|
||||
self.num_heads = num_heads
|
||||
|
||||
def forward(self, x):
|
||||
x = x.flatten(start_dim=2).permute(2, 0, 1) # NCHW -> (HW)NC
|
||||
x = torch.cat([x.mean(dim=0, keepdim=True), x], dim=0) # (HW+1)NC
|
||||
x = x + self.positional_embedding[:, None, :].to(x.dtype) # (HW+1)NC
|
||||
x, _ = F.multi_head_attention_forward(
|
||||
query=x[:1], key=x, value=x,
|
||||
embed_dim_to_check=x.shape[-1],
|
||||
num_heads=self.num_heads,
|
||||
q_proj_weight=self.q_proj.weight,
|
||||
k_proj_weight=self.k_proj.weight,
|
||||
v_proj_weight=self.v_proj.weight,
|
||||
in_proj_weight=None,
|
||||
in_proj_bias=torch.cat([self.q_proj.bias, self.k_proj.bias, self.v_proj.bias]),
|
||||
bias_k=None,
|
||||
bias_v=None,
|
||||
add_zero_attn=False,
|
||||
dropout_p=0,
|
||||
out_proj_weight=self.c_proj.weight,
|
||||
out_proj_bias=self.c_proj.bias,
|
||||
use_separate_proj_weight=True,
|
||||
training=self.training,
|
||||
need_weights=False
|
||||
)
|
||||
return x.squeeze(0)
|
||||
|
||||
|
||||
class ModifiedResNet(nn.Module):
|
||||
"""
|
||||
A ResNet class that is similar to torchvision's but contains the following changes:
|
||||
- There are now 3 "stem" convolutions as opposed to 1, with an average pool instead of a max pool.
|
||||
- Performs anti-aliasing strided convolutions, where an avgpool is prepended to convolutions with stride > 1
|
||||
- The final pooling layer is a QKV attention instead of an average pool
|
||||
"""
|
||||
|
||||
def __init__(self, layers, output_dim, heads, input_resolution=224, width=64):
|
||||
super().__init__()
|
||||
self.output_dim = output_dim
|
||||
self.input_resolution = input_resolution
|
||||
|
||||
# the 3-layer stem
|
||||
self.conv1 = nn.Conv2d(3, width // 2, kernel_size=3, stride=2, padding=1, bias=False)
|
||||
self.bn1 = nn.BatchNorm2d(width // 2)
|
||||
self.relu1 = nn.ReLU(inplace=True)
|
||||
self.conv2 = nn.Conv2d(width // 2, width // 2, kernel_size=3, padding=1, bias=False)
|
||||
self.bn2 = nn.BatchNorm2d(width // 2)
|
||||
self.relu2 = nn.ReLU(inplace=True)
|
||||
self.conv3 = nn.Conv2d(width // 2, width, kernel_size=3, padding=1, bias=False)
|
||||
self.bn3 = nn.BatchNorm2d(width)
|
||||
self.relu3 = nn.ReLU(inplace=True)
|
||||
self.avgpool = nn.AvgPool2d(2)
|
||||
|
||||
# residual layers
|
||||
self._inplanes = width # this is a *mutable* variable used during construction
|
||||
self.layer1 = self._make_layer(width, layers[0])
|
||||
self.layer2 = self._make_layer(width * 2, layers[1], stride=2)
|
||||
self.layer3 = self._make_layer(width * 4, layers[2], stride=2)
|
||||
self.layer4 = self._make_layer(width * 8, layers[3], stride=2)
|
||||
|
||||
embed_dim = width * 32 # the ResNet feature dimension
|
||||
self.attnpool = AttentionPool2d(input_resolution // 32, embed_dim, heads, output_dim)
|
||||
|
||||
def _make_layer(self, planes, blocks, stride=1):
|
||||
layers = [Bottleneck(self._inplanes, planes, stride)]
|
||||
|
||||
self._inplanes = planes * Bottleneck.expansion
|
||||
for _ in range(1, blocks):
|
||||
layers.append(Bottleneck(self._inplanes, planes))
|
||||
|
||||
return nn.Sequential(*layers)
|
||||
|
||||
def forward(self, x):
|
||||
def stem(x):
|
||||
x = self.relu1(self.bn1(self.conv1(x)))
|
||||
x = self.relu2(self.bn2(self.conv2(x)))
|
||||
x = self.relu3(self.bn3(self.conv3(x)))
|
||||
x = self.avgpool(x)
|
||||
return x
|
||||
|
||||
x = x.type(self.conv1.weight.dtype)
|
||||
x = stem(x)
|
||||
x = self.layer1(x)
|
||||
x = self.layer2(x)
|
||||
x = self.layer3(x)
|
||||
x = self.layer4(x)
|
||||
x = self.attnpool(x)
|
||||
|
||||
return x
|
||||
|
||||
|
||||
class LayerNorm(nn.LayerNorm):
|
||||
"""Subclass torch's LayerNorm to handle fp16."""
|
||||
|
||||
def forward(self, x: torch.Tensor):
|
||||
orig_type = x.dtype
|
||||
ret = super().forward(x.type(torch.float32))
|
||||
return ret.type(orig_type)
|
||||
|
||||
|
||||
class QuickGELU(nn.Module):
|
||||
def forward(self, x: torch.Tensor):
|
||||
return x * torch.sigmoid(1.702 * x)
|
||||
|
||||
|
||||
class ResidualAttentionBlock(nn.Module):
|
||||
def __init__(self, d_model: int, n_head: int, attn_mask: torch.Tensor = None):
|
||||
super().__init__()
|
||||
|
||||
self.attn = nn.MultiheadAttention(d_model, n_head)
|
||||
self.ln_1 = LayerNorm(d_model)
|
||||
self.mlp = nn.Sequential(OrderedDict([
|
||||
("c_fc", nn.Linear(d_model, d_model * 4)),
|
||||
("gelu", QuickGELU()),
|
||||
("c_proj", nn.Linear(d_model * 4, d_model))
|
||||
]))
|
||||
self.ln_2 = LayerNorm(d_model)
|
||||
self.attn_mask = attn_mask
|
||||
|
||||
def attention(self, x: torch.Tensor):
|
||||
self.attn_mask = self.attn_mask.to(dtype=x.dtype, device=x.device) if self.attn_mask is not None else None
|
||||
return self.attn(x, x, x, need_weights=False, attn_mask=self.attn_mask)[0]
|
||||
|
||||
def forward(self, x: torch.Tensor):
|
||||
x = x + self.attention(self.ln_1(x))
|
||||
x = x + self.mlp(self.ln_2(x))
|
||||
return x
|
||||
|
||||
|
||||
class Transformer(nn.Module):
|
||||
def __init__(self, width: int, layers: int, heads: int, attn_mask: torch.Tensor = None):
|
||||
super().__init__()
|
||||
self.width = width
|
||||
self.layers = layers
|
||||
self.resblocks = nn.Sequential(*[ResidualAttentionBlock(width, heads, attn_mask) for _ in range(layers)])
|
||||
|
||||
def forward(self, x: torch.Tensor):
|
||||
return self.resblocks(x)
|
||||
|
||||
|
||||
class VisionTransformer(nn.Module):
|
||||
def __init__(self, input_resolution: int, patch_size: int, width: int, layers: int, heads: int, output_dim: int):
|
||||
super().__init__()
|
||||
self.input_resolution = input_resolution
|
||||
self.output_dim = output_dim
|
||||
self.conv1 = nn.Conv2d(in_channels=3, out_channels=width, kernel_size=patch_size, stride=patch_size, bias=False)
|
||||
|
||||
scale = width ** -0.5
|
||||
self.class_embedding = nn.Parameter(scale * torch.randn(width))
|
||||
self.positional_embedding = nn.Parameter(scale * torch.randn((input_resolution // patch_size) ** 2 + 1, width))
|
||||
self.ln_pre = LayerNorm(width)
|
||||
|
||||
self.transformer = Transformer(width, layers, heads)
|
||||
|
||||
self.ln_post = LayerNorm(width)
|
||||
self.proj = nn.Parameter(scale * torch.randn(width, output_dim))
|
||||
|
||||
def forward(self, x: torch.Tensor):
|
||||
x = self.conv1(x) # shape = [*, width, grid, grid]
|
||||
x = x.reshape(x.shape[0], x.shape[1], -1) # shape = [*, width, grid ** 2]
|
||||
x = x.permute(0, 2, 1) # shape = [*, grid ** 2, width]
|
||||
x = torch.cat([self.class_embedding.to(x.dtype) + torch.zeros(x.shape[0], 1, x.shape[-1], dtype=x.dtype, device=x.device), x], dim=1) # shape = [*, grid ** 2 + 1, width]
|
||||
x = x + self.positional_embedding.to(x.dtype)
|
||||
x = self.ln_pre(x)
|
||||
|
||||
x = x.permute(1, 0, 2) # NLD -> LND
|
||||
x = self.transformer(x)
|
||||
x = x.permute(1, 0, 2) # LND -> NLD
|
||||
|
||||
x = self.ln_post(x[:, 0, :])
|
||||
|
||||
if self.proj is not None:
|
||||
x = x @ self.proj
|
||||
|
||||
return x
|
||||
|
||||
|
||||
class CLIP(nn.Module):
|
||||
def __init__(self,
|
||||
embed_dim: int,
|
||||
# vision
|
||||
image_resolution: int,
|
||||
vision_layers: Union[Tuple[int, int, int, int], int],
|
||||
vision_width: int,
|
||||
vision_patch_size: int,
|
||||
# text
|
||||
context_length: int,
|
||||
vocab_size: int,
|
||||
transformer_width: int,
|
||||
transformer_heads: int,
|
||||
transformer_layers: int,
|
||||
load_from_clip: bool
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.context_length = 248
|
||||
|
||||
if isinstance(vision_layers, (tuple, list)):
|
||||
vision_heads = vision_width * 32 // 64
|
||||
self.visual = ModifiedResNet(
|
||||
layers=vision_layers,
|
||||
output_dim=embed_dim,
|
||||
heads=vision_heads,
|
||||
input_resolution=image_resolution,
|
||||
width=vision_width
|
||||
)
|
||||
else:
|
||||
vision_heads = vision_width // 64
|
||||
self.visual = VisionTransformer(
|
||||
input_resolution=image_resolution,
|
||||
patch_size=vision_patch_size,
|
||||
width=vision_width,
|
||||
layers=vision_layers,
|
||||
heads=vision_heads,
|
||||
output_dim=embed_dim
|
||||
)
|
||||
|
||||
self.transformer = Transformer(
|
||||
width=transformer_width,
|
||||
layers=transformer_layers,
|
||||
heads=transformer_heads,
|
||||
attn_mask=self.build_attention_mask()
|
||||
)
|
||||
|
||||
self.vocab_size = vocab_size
|
||||
self.token_embedding = nn.Embedding(vocab_size, transformer_width)
|
||||
|
||||
if load_from_clip == False:
|
||||
self.positional_embedding = nn.Parameter(torch.empty(248, transformer_width))
|
||||
self.positional_embedding_res = nn.Parameter(torch.empty(248, transformer_width))
|
||||
|
||||
else:
|
||||
self.positional_embedding = nn.Parameter(torch.empty(77, transformer_width))
|
||||
|
||||
self.ln_final = LayerNorm(transformer_width)
|
||||
|
||||
self.text_projection = nn.Parameter(torch.empty(transformer_width, embed_dim))
|
||||
self.logit_scale = nn.Parameter(torch.ones([]) * np.log(1 / 0.07))
|
||||
|
||||
self.initialize_parameters()
|
||||
self.mask1 = torch.zeros([248, 1])
|
||||
self.mask1[:20, :] = 1
|
||||
self.mask2 = torch.zeros([248, 1])
|
||||
self.mask2[20:, :] = 1
|
||||
|
||||
|
||||
def initialize_parameters(self):
|
||||
nn.init.normal_(self.token_embedding.weight, std=0.02)
|
||||
nn.init.normal_(self.positional_embedding, std=0.01)
|
||||
|
||||
if isinstance(self.visual, ModifiedResNet):
|
||||
if self.visual.attnpool is not None:
|
||||
std = self.visual.attnpool.c_proj.in_features ** -0.5
|
||||
nn.init.normal_(self.visual.attnpool.q_proj.weight, std=std)
|
||||
nn.init.normal_(self.visual.attnpool.k_proj.weight, std=std)
|
||||
nn.init.normal_(self.visual.attnpool.v_proj.weight, std=std)
|
||||
nn.init.normal_(self.visual.attnpool.c_proj.weight, std=std)
|
||||
|
||||
for resnet_block in [self.visual.layer1, self.visual.layer2, self.visual.layer3, self.visual.layer4]:
|
||||
for name, param in resnet_block.named_parameters():
|
||||
if name.endswith("bn3.weight"):
|
||||
nn.init.zeros_(param)
|
||||
|
||||
proj_std = (self.transformer.width ** -0.5) * ((2 * self.transformer.layers) ** -0.5)
|
||||
attn_std = self.transformer.width ** -0.5
|
||||
fc_std = (2 * self.transformer.width) ** -0.5
|
||||
for block in self.transformer.resblocks:
|
||||
nn.init.normal_(block.attn.in_proj_weight, std=attn_std)
|
||||
nn.init.normal_(block.attn.out_proj.weight, std=proj_std)
|
||||
nn.init.normal_(block.mlp.c_fc.weight, std=fc_std)
|
||||
nn.init.normal_(block.mlp.c_proj.weight, std=proj_std)
|
||||
|
||||
if self.text_projection is not None:
|
||||
nn.init.normal_(self.text_projection, std=self.transformer.width ** -0.5)
|
||||
|
||||
def build_attention_mask(self):
|
||||
# lazily create causal attention mask, with full attention between the vision tokens
|
||||
# pytorch uses additive attention mask; fill with -inf
|
||||
mask = torch.empty(self.context_length, self.context_length)
|
||||
mask.fill_(float("-inf"))
|
||||
mask.triu_(1) # zero out the lower diagonal
|
||||
return mask
|
||||
|
||||
@property
|
||||
def dtype(self):
|
||||
# Fix: the mixclip inference.
|
||||
if hasattr(self, "visual"):
|
||||
return self.visual.conv1.weight.dtype
|
||||
else:
|
||||
return self.token_embedding.weight.dtype
|
||||
|
||||
def encode_image(self, image):
|
||||
return self.visual(image.type(self.dtype))
|
||||
|
||||
def encode_text(self, text):
|
||||
x = self.token_embedding(text).type(self.dtype) # [batch_size, n_ctx, d_model]
|
||||
|
||||
x = x + (self.positional_embedding.to(x.device) * self.mask1.to(x.device)).type(self.dtype).to(x.device) + (self.positional_embedding_res.to(x.device) * self.mask2.to(x.device)).type(self.dtype).to(x.device)
|
||||
|
||||
x = x.permute(1, 0, 2) # NLD -> LND
|
||||
x = self.transformer(x)
|
||||
x = x.permute(1, 0, 2) # LND -> NLD
|
||||
x = self.ln_final(x).type(self.dtype)
|
||||
|
||||
# x.shape = [batch_size, n_ctx, transformer.width]
|
||||
# take features from the eot embedding (eot_token is the highest number in each sequence)
|
||||
x = x[torch.arange(x.shape[0]), text.argmax(dim=-1)] @ self.text_projection
|
||||
|
||||
return x
|
||||
|
||||
def encode_text_full(self, text):
|
||||
x = self.token_embedding(text).type(self.dtype) # [batch_size, n_ctx, d_model]
|
||||
|
||||
x = x + (self.positional_embedding.to(x.device) * self.mask1.to(x.device)).type(self.dtype).to(x.device) + (self.positional_embedding_res.to(x.device) * self.mask2.to(x.device)).type(self.dtype).to(x.device)
|
||||
|
||||
x = x.permute(1, 0, 2) # NLD -> LND
|
||||
x = self.transformer(x)
|
||||
x = x.permute(1, 0, 2) # LND -> NLD
|
||||
x = self.ln_final(x).type(self.dtype)
|
||||
|
||||
# x.shape = [batch_size, n_ctx, transformer.width]
|
||||
# take features from the eot embedding (eot_token is the highest number in each sequence)
|
||||
#x = x[torch.arange(x.shape[0]), text.argmax(dim=-1)] @ self.text_projection
|
||||
|
||||
return x
|
||||
|
||||
|
||||
def forward(self, image, text):
|
||||
image_features = self.encode_image(image)
|
||||
text_features = self.encode_text(text)
|
||||
|
||||
# normalized features
|
||||
image_features = image_features / image_features.norm(dim=1, keepdim=True)
|
||||
text_features = text_features / text_features.norm(dim=1, keepdim=True)
|
||||
|
||||
# cosine similarity as logits
|
||||
logit_scale = self.logit_scale.exp()
|
||||
logits_per_image = logit_scale * image_features @ text_features.t()
|
||||
logits_per_text = logits_per_image.t()
|
||||
|
||||
# shape = [global_batch_size, global_batch_size]
|
||||
return logits_per_image, logits_per_text
|
||||
|
||||
|
||||
def convert_weights(model: nn.Module):
|
||||
"""Convert applicable model parameters to fp16"""
|
||||
|
||||
def _convert_weights_to_fp16(l):
|
||||
if isinstance(l, (nn.Conv1d, nn.Conv2d, nn.Linear)):
|
||||
l.weight.data = l.weight.data.half()
|
||||
if l.bias is not None:
|
||||
l.bias.data = l.bias.data.half()
|
||||
|
||||
if isinstance(l, nn.MultiheadAttention):
|
||||
for attr in [*[f"{s}_proj_weight" for s in ["in", "q", "k", "v"]], "in_proj_bias", "bias_k", "bias_v"]:
|
||||
tensor = getattr(l, attr)
|
||||
if tensor is not None:
|
||||
tensor.data = tensor.data.half()
|
||||
|
||||
for name in ["text_projection", "proj"]:
|
||||
if hasattr(l, name):
|
||||
attr = getattr(l, name)
|
||||
if attr is not None:
|
||||
attr.data = attr.data.half()
|
||||
|
||||
model.apply(_convert_weights_to_fp16)
|
||||
|
||||
|
||||
def build_model(state_dict: dict, load_from_clip: bool):
|
||||
vit = "visual.proj" in state_dict
|
||||
|
||||
if vit:
|
||||
vision_width = state_dict["visual.conv1.weight"].shape[0]
|
||||
vision_layers = len([k for k in state_dict.keys() if k.startswith("visual.") and k.endswith(".attn.in_proj_weight")])
|
||||
vision_patch_size = state_dict["visual.conv1.weight"].shape[-1]
|
||||
grid_size = round((state_dict["visual.positional_embedding"].shape[0] - 1) ** 0.5)
|
||||
image_resolution = vision_patch_size * grid_size
|
||||
else:
|
||||
counts: list = [len(set(k.split(".")[2] for k in state_dict if k.startswith(f"visual.layer{b}"))) for b in [1, 2, 3, 4]]
|
||||
vision_layers = tuple(counts)
|
||||
vision_width = state_dict["visual.layer1.0.conv1.weight"].shape[0]
|
||||
output_width = round((state_dict["visual.attnpool.positional_embedding"].shape[0] - 1) ** 0.5)
|
||||
vision_patch_size = None
|
||||
assert output_width ** 2 + 1 == state_dict["visual.attnpool.positional_embedding"].shape[0]
|
||||
image_resolution = output_width * 32
|
||||
|
||||
embed_dim = state_dict["text_projection"].shape[1]
|
||||
context_length = state_dict["positional_embedding"].shape[0]
|
||||
vocab_size = state_dict["token_embedding.weight"].shape[0]
|
||||
transformer_width = state_dict["ln_final.weight"].shape[0]
|
||||
transformer_heads = transformer_width // 64
|
||||
transformer_layers = len(set(k.split(".")[2] for k in state_dict if k.startswith("transformer.resblocks")))
|
||||
|
||||
model = CLIP(
|
||||
embed_dim,
|
||||
image_resolution, vision_layers, vision_width, vision_patch_size,
|
||||
context_length, vocab_size, transformer_width, transformer_heads, transformer_layers, load_from_clip
|
||||
)
|
||||
|
||||
for key in ["input_resolution", "context_length", "vocab_size"]:
|
||||
if key in state_dict:
|
||||
del state_dict[key]
|
||||
|
||||
convert_weights(model)
|
||||
model.load_state_dict(state_dict)
|
||||
return model.eval()
|
||||
@@ -0,0 +1,132 @@
|
||||
import gzip
|
||||
import html
|
||||
import os
|
||||
from functools import lru_cache
|
||||
|
||||
import ftfy
|
||||
import regex as re
|
||||
|
||||
|
||||
@lru_cache()
|
||||
def default_bpe():
|
||||
return os.path.join(os.path.dirname(os.path.abspath(__file__)), "bpe_simple_vocab_16e6.txt.gz")
|
||||
|
||||
|
||||
@lru_cache()
|
||||
def bytes_to_unicode():
|
||||
"""
|
||||
Returns list of utf-8 byte and a corresponding list of unicode strings.
|
||||
The reversible bpe codes work on unicode strings.
|
||||
This means you need a large # of unicode characters in your vocab if you want to avoid UNKs.
|
||||
When you're at something like a 10B token dataset you end up needing around 5K for decent coverage.
|
||||
This is a signficant percentage of your normal, say, 32K bpe vocab.
|
||||
To avoid that, we want lookup tables between utf-8 bytes and unicode strings.
|
||||
And avoids mapping to whitespace/control characters the bpe code barfs on.
|
||||
"""
|
||||
bs = list(range(ord("!"), ord("~")+1))+list(range(ord("¡"), ord("¬")+1))+list(range(ord("®"), ord("ÿ")+1))
|
||||
cs = bs[:]
|
||||
n = 0
|
||||
for b in range(2**8):
|
||||
if b not in bs:
|
||||
bs.append(b)
|
||||
cs.append(2**8+n)
|
||||
n += 1
|
||||
cs = [chr(n) for n in cs]
|
||||
return dict(zip(bs, cs))
|
||||
|
||||
|
||||
def get_pairs(word):
|
||||
"""Return set of symbol pairs in a word.
|
||||
Word is represented as tuple of symbols (symbols being variable-length strings).
|
||||
"""
|
||||
pairs = set()
|
||||
prev_char = word[0]
|
||||
for char in word[1:]:
|
||||
pairs.add((prev_char, char))
|
||||
prev_char = char
|
||||
return pairs
|
||||
|
||||
|
||||
def basic_clean(text):
|
||||
text = ftfy.fix_text(text)
|
||||
text = html.unescape(html.unescape(text))
|
||||
return text.strip()
|
||||
|
||||
|
||||
def whitespace_clean(text):
|
||||
text = re.sub(r'\s+', ' ', text)
|
||||
text = text.strip()
|
||||
return text
|
||||
|
||||
|
||||
class SimpleTokenizer(object):
|
||||
def __init__(self, bpe_path: str = default_bpe()):
|
||||
self.byte_encoder = bytes_to_unicode()
|
||||
self.byte_decoder = {v: k for k, v in self.byte_encoder.items()}
|
||||
merges = gzip.open(bpe_path).read().decode("utf-8").split('\n')
|
||||
merges = merges[1:49152-256-2+1]
|
||||
merges = [tuple(merge.split()) for merge in merges]
|
||||
vocab = list(bytes_to_unicode().values())
|
||||
vocab = vocab + [v+'</w>' for v in vocab]
|
||||
for merge in merges:
|
||||
vocab.append(''.join(merge))
|
||||
vocab.extend(['<|startoftext|>', '<|endoftext|>'])
|
||||
self.encoder = dict(zip(vocab, range(len(vocab))))
|
||||
self.decoder = {v: k for k, v in self.encoder.items()}
|
||||
self.bpe_ranks = dict(zip(merges, range(len(merges))))
|
||||
self.cache = {'<|startoftext|>': '<|startoftext|>', '<|endoftext|>': '<|endoftext|>'}
|
||||
self.pat = re.compile(r"""<\|startoftext\|>|<\|endoftext\|>|'s|'t|'re|'ve|'m|'ll|'d|[\p{L}]+|[\p{N}]|[^\s\p{L}\p{N}]+""", re.IGNORECASE)
|
||||
|
||||
def bpe(self, token):
|
||||
if token in self.cache:
|
||||
return self.cache[token]
|
||||
word = tuple(token[:-1]) + ( token[-1] + '</w>',)
|
||||
pairs = get_pairs(word)
|
||||
|
||||
if not pairs:
|
||||
return token+'</w>'
|
||||
|
||||
while True:
|
||||
bigram = min(pairs, key = lambda pair: self.bpe_ranks.get(pair, float('inf')))
|
||||
if bigram not in self.bpe_ranks:
|
||||
break
|
||||
first, second = bigram
|
||||
new_word = []
|
||||
i = 0
|
||||
while i < len(word):
|
||||
try:
|
||||
j = word.index(first, i)
|
||||
new_word.extend(word[i:j])
|
||||
i = j
|
||||
except:
|
||||
new_word.extend(word[i:])
|
||||
break
|
||||
|
||||
if word[i] == first and i < len(word)-1 and word[i+1] == second:
|
||||
new_word.append(first+second)
|
||||
i += 2
|
||||
else:
|
||||
new_word.append(word[i])
|
||||
i += 1
|
||||
new_word = tuple(new_word)
|
||||
word = new_word
|
||||
if len(word) == 1:
|
||||
break
|
||||
else:
|
||||
pairs = get_pairs(word)
|
||||
word = ' '.join(word)
|
||||
self.cache[token] = word
|
||||
return word
|
||||
|
||||
def encode(self, text):
|
||||
bpe_tokens = []
|
||||
text = whitespace_clean(basic_clean(text)).lower()
|
||||
for token in re.findall(self.pat, text):
|
||||
token = ''.join(self.byte_encoder[b] for b in token.encode('utf-8'))
|
||||
bpe_tokens.extend(self.encoder[bpe_token] for bpe_token in self.bpe(token).split(' '))
|
||||
return bpe_tokens
|
||||
|
||||
def decode(self, tokens):
|
||||
text = ''.join([self.decoder[token] for token in tokens])
|
||||
text = bytearray([self.byte_decoder[c] for c in text]).decode('utf-8', errors="replace").replace('</w>', ' ')
|
||||
return text
|
||||
@@ -0,0 +1,127 @@
|
||||
# Borrowed from https://github.com/discus0434/aesthetic-predictor-v2-5/blob/3125a9e/src/aesthetic_predictor_v2_5/siglip_v2_5.py.
|
||||
import os
|
||||
from collections import OrderedDict
|
||||
from os import PathLike
|
||||
from typing import Final
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from transformers import (
|
||||
SiglipImageProcessor,
|
||||
SiglipVisionConfig,
|
||||
SiglipVisionModel,
|
||||
logging,
|
||||
)
|
||||
from transformers.image_processing_utils import BatchFeature
|
||||
from transformers.modeling_outputs import ImageClassifierOutputWithNoAttention
|
||||
|
||||
logging.set_verbosity_error()
|
||||
|
||||
URL: Final[str] = (
|
||||
"https://github.com/discus0434/aesthetic-predictor-v2-5/raw/main/models/aesthetic_predictor_v2_5.pth"
|
||||
)
|
||||
|
||||
|
||||
class AestheticPredictorV2_5Head(nn.Module):
|
||||
def __init__(self, config: SiglipVisionConfig) -> None:
|
||||
super().__init__()
|
||||
self.scoring_head = nn.Sequential(
|
||||
nn.Linear(config.hidden_size, 1024),
|
||||
nn.Dropout(0.5),
|
||||
nn.Linear(1024, 128),
|
||||
nn.Dropout(0.5),
|
||||
nn.Linear(128, 64),
|
||||
nn.Dropout(0.5),
|
||||
nn.Linear(64, 16),
|
||||
nn.Dropout(0.2),
|
||||
nn.Linear(16, 1),
|
||||
)
|
||||
|
||||
def forward(self, image_embeds: torch.Tensor) -> torch.Tensor:
|
||||
return self.scoring_head(image_embeds)
|
||||
|
||||
|
||||
class AestheticPredictorV2_5Model(SiglipVisionModel):
|
||||
PATCH_SIZE = 14
|
||||
|
||||
def __init__(self, config: SiglipVisionConfig, *args, **kwargs) -> None:
|
||||
super().__init__(config, *args, **kwargs)
|
||||
self.layers = AestheticPredictorV2_5Head(config)
|
||||
self.post_init()
|
||||
|
||||
def forward(
|
||||
self,
|
||||
pixel_values: torch.FloatTensor | None = None,
|
||||
labels: torch.Tensor | None = None,
|
||||
return_dict: bool | None = None,
|
||||
) -> tuple | ImageClassifierOutputWithNoAttention:
|
||||
return_dict = (
|
||||
return_dict if return_dict is not None else self.config.use_return_dict
|
||||
)
|
||||
|
||||
outputs = super().forward(
|
||||
pixel_values=pixel_values,
|
||||
return_dict=return_dict,
|
||||
)
|
||||
image_embeds = outputs.pooler_output
|
||||
image_embeds_norm = image_embeds / image_embeds.norm(dim=-1, keepdim=True)
|
||||
prediction = self.layers(image_embeds_norm)
|
||||
|
||||
loss = None
|
||||
if labels is not None:
|
||||
loss_fct = nn.MSELoss()
|
||||
loss = loss_fct()
|
||||
|
||||
if not return_dict:
|
||||
return (loss, prediction, image_embeds)
|
||||
|
||||
return ImageClassifierOutputWithNoAttention(
|
||||
loss=loss,
|
||||
logits=prediction,
|
||||
hidden_states=image_embeds,
|
||||
)
|
||||
|
||||
|
||||
class AestheticPredictorV2_5Processor(SiglipImageProcessor):
|
||||
def __init__(self, *args, **kwargs) -> None:
|
||||
super().__init__(*args, **kwargs)
|
||||
|
||||
def __call__(self, *args, **kwargs) -> BatchFeature:
|
||||
return super().__call__(*args, **kwargs)
|
||||
|
||||
@classmethod
|
||||
def from_pretrained(
|
||||
self,
|
||||
pretrained_model_name_or_path: str
|
||||
| PathLike = "google/siglip-so400m-patch14-384",
|
||||
*args,
|
||||
**kwargs,
|
||||
) -> "AestheticPredictorV2_5Processor":
|
||||
return super().from_pretrained(pretrained_model_name_or_path, *args, **kwargs)
|
||||
|
||||
|
||||
def convert_v2_5_from_siglip(
|
||||
predictor_name_or_path: str | PathLike | None = None,
|
||||
encoder_model_name: str = "google/siglip-so400m-patch14-384",
|
||||
*args,
|
||||
**kwargs,
|
||||
) -> tuple[AestheticPredictorV2_5Model, AestheticPredictorV2_5Processor]:
|
||||
model = AestheticPredictorV2_5Model.from_pretrained(
|
||||
encoder_model_name, *args, **kwargs
|
||||
)
|
||||
|
||||
processor = AestheticPredictorV2_5Processor.from_pretrained(
|
||||
encoder_model_name, *args, **kwargs
|
||||
)
|
||||
|
||||
if predictor_name_or_path is None or not os.path.exists(predictor_name_or_path):
|
||||
state_dict = torch.hub.load_state_dict_from_url(URL, map_location="cpu")
|
||||
else:
|
||||
state_dict = torch.load(predictor_name_or_path, map_location="cpu")
|
||||
|
||||
assert isinstance(state_dict, OrderedDict)
|
||||
|
||||
model.layers.load_state_dict(state_dict)
|
||||
model.eval()
|
||||
|
||||
return model, processor
|
||||
@@ -0,0 +1,2 @@
|
||||
# ViCLIP
|
||||
Codes in this directory are borrowed from https://github.com/OpenGVLab/InternVideo/tree/73271ba/Data/InternVid/viclip.
|
||||
@@ -0,0 +1,72 @@
|
||||
from .simple_tokenizer import SimpleTokenizer as _Tokenizer
|
||||
from .viclip import ViCLIP
|
||||
import torch
|
||||
import numpy as np
|
||||
import cv2
|
||||
import os
|
||||
|
||||
|
||||
def get_viclip(size='l',
|
||||
pretrain=os.path.join(os.path.dirname(os.path.abspath(__file__)), "ViClip-InternVid-10M-FLT.pth")):
|
||||
|
||||
tokenizer = _Tokenizer()
|
||||
vclip = ViCLIP(tokenizer=tokenizer, size=size, pretrain=pretrain)
|
||||
m = {'viclip':vclip, 'tokenizer':tokenizer}
|
||||
|
||||
return m
|
||||
|
||||
def get_text_feat_dict(texts, clip, tokenizer, text_feat_d={}):
|
||||
for t in texts:
|
||||
feat = clip.get_text_features(t, tokenizer, text_feat_d)
|
||||
text_feat_d[t] = feat
|
||||
return text_feat_d
|
||||
|
||||
def get_vid_feat(frames, clip):
|
||||
return clip.get_vid_features(frames)
|
||||
|
||||
def _frame_from_video(video):
|
||||
while video.isOpened():
|
||||
success, frame = video.read()
|
||||
if success:
|
||||
yield frame
|
||||
else:
|
||||
break
|
||||
|
||||
v_mean = np.array([0.485, 0.456, 0.406]).reshape(1,1,3)
|
||||
v_std = np.array([0.229, 0.224, 0.225]).reshape(1,1,3)
|
||||
def normalize(data):
|
||||
return (data/255.0-v_mean)/v_std
|
||||
|
||||
def frames2tensor(vid_list, fnum=8, target_size=(224, 224), device=torch.device('cuda')):
|
||||
assert(len(vid_list) >= fnum)
|
||||
step = len(vid_list) // fnum
|
||||
vid_list = vid_list[::step][:fnum]
|
||||
vid_list = [cv2.resize(x[:,:,::-1], target_size) for x in vid_list]
|
||||
vid_tube = [np.expand_dims(normalize(x), axis=(0, 1)) for x in vid_list]
|
||||
vid_tube = np.concatenate(vid_tube, axis=1)
|
||||
vid_tube = np.transpose(vid_tube, (0, 1, 4, 2, 3))
|
||||
vid_tube = torch.from_numpy(vid_tube).to(device, non_blocking=True).float()
|
||||
return vid_tube
|
||||
|
||||
def retrieve_text(frames,
|
||||
texts,
|
||||
models={'viclip':None,
|
||||
'tokenizer':None},
|
||||
topk=5,
|
||||
device=torch.device('cuda')):
|
||||
# clip, tokenizer = get_clip(name, model_cfg['size'], model_cfg['pretrained'], model_cfg['reload'])
|
||||
assert(type(models)==dict and models['viclip'] is not None and models['tokenizer'] is not None)
|
||||
clip, tokenizer = models['viclip'], models['tokenizer']
|
||||
clip = clip.to(device)
|
||||
frames_tensor = frames2tensor(frames, device=device)
|
||||
vid_feat = get_vid_feat(frames_tensor, clip)
|
||||
|
||||
text_feat_d = {}
|
||||
text_feat_d = get_text_feat_dict(texts, clip, tokenizer, text_feat_d)
|
||||
text_feats = [text_feat_d[t] for t in texts]
|
||||
text_feats_tensor = torch.cat(text_feats, 0)
|
||||
|
||||
probs, idxs = clip.get_predict_label(vid_feat, text_feats_tensor, top=topk)
|
||||
|
||||
ret_texts = [texts[i] for i in idxs.numpy()[0].tolist()]
|
||||
return ret_texts, probs.numpy()[0]
|
||||
Binary file not shown.
@@ -0,0 +1,135 @@
|
||||
import gzip
|
||||
import html
|
||||
import os
|
||||
from functools import lru_cache
|
||||
|
||||
import ftfy
|
||||
import regex as re
|
||||
|
||||
|
||||
@lru_cache()
|
||||
def default_bpe():
|
||||
return os.path.join(os.path.dirname(os.path.abspath(__file__)), "bpe_simple_vocab_16e6.txt.gz")
|
||||
# @lru_cache()
|
||||
# def default_bpe():
|
||||
# return "bpe_simple_vocab_16e6.txt.gz"
|
||||
|
||||
|
||||
@lru_cache()
|
||||
def bytes_to_unicode():
|
||||
"""
|
||||
Returns list of utf-8 byte and a corresponding list of unicode strings.
|
||||
The reversible bpe codes work on unicode strings.
|
||||
This means you need a large # of unicode characters in your vocab if you want to avoid UNKs.
|
||||
When you're at something like a 10B token dataset you end up needing around 5K for decent coverage.
|
||||
This is a signficant percentage of your normal, say, 32K bpe vocab.
|
||||
To avoid that, we want lookup tables between utf-8 bytes and unicode strings.
|
||||
And avoids mapping to whitespace/control characters the bpe code barfs on.
|
||||
"""
|
||||
bs = list(range(ord("!"), ord("~")+1))+list(range(ord("¡"), ord("¬")+1))+list(range(ord("®"), ord("ÿ")+1))
|
||||
cs = bs[:]
|
||||
n = 0
|
||||
for b in range(2**8):
|
||||
if b not in bs:
|
||||
bs.append(b)
|
||||
cs.append(2**8+n)
|
||||
n += 1
|
||||
cs = [chr(n) for n in cs]
|
||||
return dict(zip(bs, cs))
|
||||
|
||||
|
||||
def get_pairs(word):
|
||||
"""Return set of symbol pairs in a word.
|
||||
Word is represented as tuple of symbols (symbols being variable-length strings).
|
||||
"""
|
||||
pairs = set()
|
||||
prev_char = word[0]
|
||||
for char in word[1:]:
|
||||
pairs.add((prev_char, char))
|
||||
prev_char = char
|
||||
return pairs
|
||||
|
||||
|
||||
def basic_clean(text):
|
||||
text = ftfy.fix_text(text)
|
||||
text = html.unescape(html.unescape(text))
|
||||
return text.strip()
|
||||
|
||||
|
||||
def whitespace_clean(text):
|
||||
text = re.sub(r'\s+', ' ', text)
|
||||
text = text.strip()
|
||||
return text
|
||||
|
||||
|
||||
class SimpleTokenizer(object):
|
||||
def __init__(self, bpe_path: str = default_bpe()):
|
||||
self.byte_encoder = bytes_to_unicode()
|
||||
self.byte_decoder = {v: k for k, v in self.byte_encoder.items()}
|
||||
merges = gzip.open(bpe_path).read().decode("utf-8").split('\n')
|
||||
merges = merges[1:49152-256-2+1]
|
||||
merges = [tuple(merge.split()) for merge in merges]
|
||||
vocab = list(bytes_to_unicode().values())
|
||||
vocab = vocab + [v+'</w>' for v in vocab]
|
||||
for merge in merges:
|
||||
vocab.append(''.join(merge))
|
||||
vocab.extend(['<|startoftext|>', '<|endoftext|>'])
|
||||
self.encoder = dict(zip(vocab, range(len(vocab))))
|
||||
self.decoder = {v: k for k, v in self.encoder.items()}
|
||||
self.bpe_ranks = dict(zip(merges, range(len(merges))))
|
||||
self.cache = {'<|startoftext|>': '<|startoftext|>', '<|endoftext|>': '<|endoftext|>'}
|
||||
self.pat = re.compile(r"""<\|startoftext\|>|<\|endoftext\|>|'s|'t|'re|'ve|'m|'ll|'d|[\p{L}]+|[\p{N}]|[^\s\p{L}\p{N}]+""", re.IGNORECASE)
|
||||
|
||||
def bpe(self, token):
|
||||
if token in self.cache:
|
||||
return self.cache[token]
|
||||
word = tuple(token[:-1]) + ( token[-1] + '</w>',)
|
||||
pairs = get_pairs(word)
|
||||
|
||||
if not pairs:
|
||||
return token+'</w>'
|
||||
|
||||
while True:
|
||||
bigram = min(pairs, key = lambda pair: self.bpe_ranks.get(pair, float('inf')))
|
||||
if bigram not in self.bpe_ranks:
|
||||
break
|
||||
first, second = bigram
|
||||
new_word = []
|
||||
i = 0
|
||||
while i < len(word):
|
||||
try:
|
||||
j = word.index(first, i)
|
||||
new_word.extend(word[i:j])
|
||||
i = j
|
||||
except:
|
||||
new_word.extend(word[i:])
|
||||
break
|
||||
|
||||
if word[i] == first and i < len(word)-1 and word[i+1] == second:
|
||||
new_word.append(first+second)
|
||||
i += 2
|
||||
else:
|
||||
new_word.append(word[i])
|
||||
i += 1
|
||||
new_word = tuple(new_word)
|
||||
word = new_word
|
||||
if len(word) == 1:
|
||||
break
|
||||
else:
|
||||
pairs = get_pairs(word)
|
||||
word = ' '.join(word)
|
||||
self.cache[token] = word
|
||||
return word
|
||||
|
||||
def encode(self, text):
|
||||
bpe_tokens = []
|
||||
text = whitespace_clean(basic_clean(text)).lower()
|
||||
for token in re.findall(self.pat, text):
|
||||
token = ''.join(self.byte_encoder[b] for b in token.encode('utf-8'))
|
||||
bpe_tokens.extend(self.encoder[bpe_token] for bpe_token in self.bpe(token).split(' '))
|
||||
return bpe_tokens
|
||||
|
||||
def decode(self, tokens):
|
||||
text = ''.join([self.decoder[token] for token in tokens])
|
||||
text = bytearray([self.byte_decoder[c] for c in text]).decode('utf-8', errors="replace").replace('</w>', ' ')
|
||||
return text
|
||||
@@ -0,0 +1,262 @@
|
||||
import os
|
||||
import logging
|
||||
|
||||
import torch
|
||||
from einops import rearrange
|
||||
from torch import nn
|
||||
import math
|
||||
|
||||
# from .criterions import VTC_VTM_Loss
|
||||
from .simple_tokenizer import SimpleTokenizer as _Tokenizer
|
||||
from .viclip_vision import clip_joint_l14, clip_joint_b16
|
||||
from .viclip_text import clip_text_l14, clip_text_b16
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class ViCLIP(nn.Module):
|
||||
"""docstring for ViCLIP"""
|
||||
|
||||
def __init__(self,
|
||||
tokenizer=None,
|
||||
size='l',
|
||||
pretrain=os.path.join(os.path.dirname(os.path.abspath(__file__)), "ViClip-InternVid-10M-FLT.pth"),
|
||||
freeze_text=True):
|
||||
super(ViCLIP, self).__init__()
|
||||
if tokenizer:
|
||||
self.tokenizer = tokenizer
|
||||
else:
|
||||
self.tokenizer = _Tokenizer()
|
||||
self.max_txt_l = 32
|
||||
|
||||
if size.lower() == 'l':
|
||||
self.vision_encoder_name = 'vit_l14'
|
||||
elif size.lower() == 'b':
|
||||
self.vision_encoder_name = 'vit_b16'
|
||||
else:
|
||||
raise NotImplementedError(f"Size {size} not implemented")
|
||||
|
||||
self.vision_encoder_pretrained = False
|
||||
self.inputs_image_res = 224
|
||||
self.vision_encoder_kernel_size = 1
|
||||
self.vision_encoder_center = True
|
||||
self.video_input_num_frames = 8
|
||||
self.vision_encoder_drop_path_rate = 0.1
|
||||
self.vision_encoder_checkpoint_num = 24
|
||||
self.is_pretrain = pretrain
|
||||
self.vision_width = 1024
|
||||
self.text_width = 768
|
||||
self.embed_dim = 768
|
||||
self.masking_prob = 0.9
|
||||
|
||||
if size.lower() == 'l':
|
||||
self.text_encoder_name = 'vit_l14'
|
||||
elif size.lower() == 'b':
|
||||
self.text_encoder_name = 'vit_b16'
|
||||
else:
|
||||
raise NotImplementedError(f"Size {size} not implemented")
|
||||
|
||||
self.text_encoder_pretrained = False#'bert-base-uncased'
|
||||
self.text_encoder_d_model = 768
|
||||
|
||||
self.text_encoder_vocab_size = 49408
|
||||
|
||||
# create modules.
|
||||
self.vision_encoder = self.build_vision_encoder()
|
||||
self.text_encoder = self.build_text_encoder()
|
||||
|
||||
self.temp = nn.parameter.Parameter(torch.ones([]) * 1 / 100.0)
|
||||
self.temp_min = 1 / 100.0
|
||||
|
||||
if pretrain:
|
||||
logger.info(f"Load pretrained weights from {pretrain}")
|
||||
state_dict = torch.load(pretrain, map_location='cpu')['model']
|
||||
self.load_state_dict(state_dict)
|
||||
|
||||
# Freeze weights
|
||||
if freeze_text:
|
||||
self.freeze_text()
|
||||
|
||||
|
||||
def freeze_text(self):
|
||||
"""freeze text encoder"""
|
||||
for p in self.text_encoder.parameters():
|
||||
p.requires_grad = False
|
||||
|
||||
def no_weight_decay(self):
|
||||
ret = {"temp"}
|
||||
ret.update(
|
||||
{"vision_encoder." + k for k in self.vision_encoder.no_weight_decay()}
|
||||
)
|
||||
ret.update(
|
||||
{"text_encoder." + k for k in self.text_encoder.no_weight_decay()}
|
||||
)
|
||||
|
||||
return ret
|
||||
|
||||
def forward(self, image, text, raw_text, idx, log_generation=None, return_sims=False):
|
||||
"""forward and calculate loss.
|
||||
|
||||
Args:
|
||||
image (torch.Tensor): The input images. Shape: [B,T,C,H,W].
|
||||
text (dict): TODO
|
||||
idx (torch.Tensor): TODO
|
||||
|
||||
Returns: TODO
|
||||
|
||||
"""
|
||||
self.clip_contrastive_temperature()
|
||||
|
||||
vision_embeds = self.encode_vision(image)
|
||||
text_embeds = self.encode_text(raw_text)
|
||||
if return_sims:
|
||||
sims = torch.nn.functional.normalize(vision_embeds, dim=-1) @ \
|
||||
torch.nn.functional.normalize(text_embeds, dim=-1).transpose(0, 1)
|
||||
return sims
|
||||
|
||||
# calculate loss
|
||||
|
||||
## VTC loss
|
||||
loss_vtc = self.clip_loss.vtc_loss(
|
||||
vision_embeds, text_embeds, idx, self.temp, all_gather=True
|
||||
)
|
||||
|
||||
return dict(
|
||||
loss_vtc=loss_vtc,
|
||||
)
|
||||
|
||||
def encode_vision(self, image, test=False):
|
||||
"""encode image / videos as features.
|
||||
|
||||
Args:
|
||||
image (torch.Tensor): The input images.
|
||||
test (bool): Whether testing.
|
||||
|
||||
Returns: tuple.
|
||||
- vision_embeds (torch.Tensor): The features of all patches. Shape: [B,T,L,C].
|
||||
- pooled_vision_embeds (torch.Tensor): The pooled features. Shape: [B,T,C].
|
||||
|
||||
"""
|
||||
if image.ndim == 5:
|
||||
image = image.permute(0, 2, 1, 3, 4).contiguous()
|
||||
else:
|
||||
image = image.unsqueeze(2)
|
||||
|
||||
if not test and self.masking_prob > 0.0:
|
||||
return self.vision_encoder(
|
||||
image, masking_prob=self.masking_prob
|
||||
)
|
||||
|
||||
return self.vision_encoder(image)
|
||||
|
||||
def encode_text(self, text):
|
||||
"""encode text.
|
||||
Args:
|
||||
text (dict): The output of huggingface's `PreTrainedTokenizer`. contains keys:
|
||||
- input_ids (torch.Tensor): Token ids to be fed to a model. Shape: [B,L].
|
||||
- attention_mask (torch.Tensor): The mask indicate padded tokens. Shape: [B,L]. 0 is padded token.
|
||||
- other keys refer to "https://huggingface.co/docs/transformers/v4.21.2/en/main_classes/tokenizer#transformers.PreTrainedTokenizer.__call__".
|
||||
Returns: tuple.
|
||||
- text_embeds (torch.Tensor): The features of all tokens. Shape: [B,L,C].
|
||||
- pooled_text_embeds (torch.Tensor): The pooled features. Shape: [B,C].
|
||||
|
||||
"""
|
||||
device = next(self.text_encoder.parameters()).device
|
||||
text = self.text_encoder.tokenize(
|
||||
text, context_length=self.max_txt_l
|
||||
).to(device)
|
||||
text_embeds = self.text_encoder(text)
|
||||
return text_embeds
|
||||
|
||||
@torch.no_grad()
|
||||
def clip_contrastive_temperature(self, min_val=0.001, max_val=0.5):
|
||||
"""Seems only used during pre-training"""
|
||||
self.temp.clamp_(min=self.temp_min)
|
||||
|
||||
def build_vision_encoder(self):
|
||||
"""build vision encoder
|
||||
Returns: (vision_encoder, vision_layernorm). Each is a `nn.Module`.
|
||||
|
||||
"""
|
||||
encoder_name = self.vision_encoder_name
|
||||
if encoder_name == "vit_l14":
|
||||
vision_encoder = clip_joint_l14(
|
||||
pretrained=self.vision_encoder_pretrained,
|
||||
input_resolution=self.inputs_image_res,
|
||||
kernel_size=self.vision_encoder_kernel_size,
|
||||
center=self.vision_encoder_center,
|
||||
num_frames=self.video_input_num_frames,
|
||||
drop_path=self.vision_encoder_drop_path_rate,
|
||||
checkpoint_num=self.vision_encoder_checkpoint_num,
|
||||
)
|
||||
elif encoder_name == "vit_b16":
|
||||
vision_encoder = clip_joint_b16(
|
||||
pretrained=self.vision_encoder_pretrained,
|
||||
input_resolution=self.inputs_image_res,
|
||||
kernel_size=self.vision_encoder_kernel_size,
|
||||
center=self.vision_encoder_center,
|
||||
num_frames=self.video_input_num_frames,
|
||||
drop_path=self.vision_encoder_drop_path_rate,
|
||||
checkpoint_num=self.vision_encoder_checkpoint_num,
|
||||
)
|
||||
else:
|
||||
raise NotImplementedError(f"Not implemented: {encoder_name}")
|
||||
|
||||
return vision_encoder
|
||||
|
||||
def build_text_encoder(self):
|
||||
"""build text_encoder and possiblly video-to-text multimodal fusion encoder.
|
||||
Returns: nn.Module. The text encoder
|
||||
|
||||
"""
|
||||
encoder_name = self.text_encoder_name
|
||||
|
||||
if encoder_name == "vit_l14":
|
||||
text_encoder = clip_text_l14(
|
||||
pretrained=self.text_encoder_pretrained,
|
||||
context_length=self.max_txt_l,
|
||||
vocab_size=self.text_encoder_vocab_size,
|
||||
checkpoint_num=0,
|
||||
)
|
||||
elif encoder_name == "vit_b16":
|
||||
text_encoder = clip_text_b16(
|
||||
pretrained=self.text_encoder_pretrained,
|
||||
context_length=self.max_txt_l,
|
||||
vocab_size=self.text_encoder_vocab_size,
|
||||
checkpoint_num=0,
|
||||
)
|
||||
else:
|
||||
raise NotImplementedError(f"Not implemented: {encoder_name}")
|
||||
|
||||
return text_encoder
|
||||
|
||||
def get_text_encoder(self):
|
||||
"""get text encoder, used for text and cross-modal encoding"""
|
||||
encoder = self.text_encoder
|
||||
return encoder.bert if hasattr(encoder, "bert") else encoder
|
||||
|
||||
def get_text_features(self, input_text, tokenizer, text_feature_dict={}):
|
||||
if input_text in text_feature_dict:
|
||||
return text_feature_dict[input_text]
|
||||
text_template= f"{input_text}"
|
||||
with torch.no_grad():
|
||||
# text_token = tokenizer.encode(text_template).cuda()
|
||||
text_features = self.encode_text(text_template).float()
|
||||
text_features /= text_features.norm(dim=-1, keepdim=True)
|
||||
text_feature_dict[input_text] = text_features
|
||||
return text_features
|
||||
|
||||
def get_vid_features(self, input_frames):
|
||||
with torch.no_grad():
|
||||
clip_feat = self.encode_vision(input_frames,test=True).float()
|
||||
clip_feat /= clip_feat.norm(dim=-1, keepdim=True)
|
||||
return clip_feat
|
||||
|
||||
def get_predict_label(self, clip_feature, text_feats_tensor, top=5):
|
||||
label_probs = (100.0 * clip_feature @ text_feats_tensor.T).softmax(dim=-1)
|
||||
top_probs, top_labels = label_probs.cpu().topk(top, dim=-1)
|
||||
return top_probs, top_labels
|
||||
|
||||
|
||||
if __name__ =="__main__":
|
||||
tokenizer = _Tokenizer()
|
||||
@@ -0,0 +1,297 @@
|
||||
import os
|
||||
import logging
|
||||
from collections import OrderedDict
|
||||
from pkg_resources import packaging
|
||||
from .simple_tokenizer import SimpleTokenizer as _Tokenizer
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from torch import nn
|
||||
import torch.utils.checkpoint as checkpoint
|
||||
import functools
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
# On P1, model extracted from https://huggingface.co/laion/CLIP-ViT-L-14-DataComp.XL-s13B-b90K
|
||||
MODEL_PATH = 'https://huggingface.co/laion'
|
||||
_MODELS = {
|
||||
"ViT-L/14": os.path.join(MODEL_PATH, "CLIP-ViT-L-14-DataComp.XL-s13B-b90K", "vit_l14_text.pth"),
|
||||
"ViT-B/16": os.path.join(MODEL_PATH, "CLIP-ViT-B-16-DataComp.XL-s13B-b90K", "vit_b16_text.pth"),
|
||||
}
|
||||
|
||||
|
||||
class LayerNorm(nn.LayerNorm):
|
||||
"""Subclass torch's LayerNorm to handle fp16."""
|
||||
|
||||
def forward(self, x: torch.Tensor):
|
||||
orig_type = x.dtype
|
||||
ret = super().forward(x.type(torch.float32))
|
||||
return ret.type(orig_type)
|
||||
|
||||
|
||||
class QuickGELU(nn.Module):
|
||||
def forward(self, x: torch.Tensor):
|
||||
return x * torch.sigmoid(1.702 * x)
|
||||
|
||||
|
||||
class ResidualAttentionBlock(nn.Module):
|
||||
def __init__(self, d_model: int, n_head: int, attn_mask: torch.Tensor = None):
|
||||
super().__init__()
|
||||
|
||||
self.attn = nn.MultiheadAttention(d_model, n_head)
|
||||
self.ln_1 = LayerNorm(d_model)
|
||||
self.mlp = nn.Sequential(OrderedDict([
|
||||
("c_fc", nn.Linear(d_model, d_model * 4)),
|
||||
("gelu", QuickGELU()),
|
||||
("c_proj", nn.Linear(d_model * 4, d_model))
|
||||
]))
|
||||
self.ln_2 = LayerNorm(d_model)
|
||||
self.attn_mask = attn_mask
|
||||
|
||||
def attention(self, x: torch.Tensor):
|
||||
self.attn_mask = self.attn_mask.to(dtype=x.dtype, device=x.device) if self.attn_mask is not None else None
|
||||
return self.attn(x, x, x, need_weights=False, attn_mask=self.attn_mask)[0]
|
||||
|
||||
def forward(self, x: torch.Tensor):
|
||||
x = x + self.attention(self.ln_1(x))
|
||||
x = x + self.mlp(self.ln_2(x))
|
||||
return x
|
||||
|
||||
|
||||
class Transformer(nn.Module):
|
||||
def __init__(self, width: int, layers: int, heads: int, attn_mask: torch.Tensor = None,
|
||||
checkpoint_num: int = 0):
|
||||
super().__init__()
|
||||
self.width = width
|
||||
self.layers = layers
|
||||
self.resblocks = nn.Sequential(*[ResidualAttentionBlock(width, heads, attn_mask) for _ in range(layers)])
|
||||
|
||||
self.checkpoint_num = checkpoint_num
|
||||
|
||||
def forward(self, x: torch.Tensor):
|
||||
if self.checkpoint_num > 0:
|
||||
segments = min(self.checkpoint_num, len(self.resblocks))
|
||||
return checkpoint.checkpoint_sequential(self.resblocks, segments, x)
|
||||
else:
|
||||
return self.resblocks(x)
|
||||
|
||||
|
||||
class CLIP_TEXT(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
embed_dim: int,
|
||||
context_length: int,
|
||||
vocab_size: int,
|
||||
transformer_width: int,
|
||||
transformer_heads: int,
|
||||
transformer_layers: int,
|
||||
checkpoint_num: int,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.context_length = context_length
|
||||
self._tokenizer = _Tokenizer()
|
||||
|
||||
self.transformer = Transformer(
|
||||
width=transformer_width,
|
||||
layers=transformer_layers,
|
||||
heads=transformer_heads,
|
||||
attn_mask=self.build_attention_mask(),
|
||||
checkpoint_num=checkpoint_num,
|
||||
)
|
||||
|
||||
self.vocab_size = vocab_size
|
||||
self.token_embedding = nn.Embedding(vocab_size, transformer_width)
|
||||
self.positional_embedding = nn.Parameter(torch.empty(self.context_length, transformer_width))
|
||||
self.ln_final = LayerNorm(transformer_width)
|
||||
|
||||
self.text_projection = nn.Parameter(torch.empty(transformer_width, embed_dim))
|
||||
|
||||
def no_weight_decay(self):
|
||||
return {'token_embedding', 'positional_embedding'}
|
||||
|
||||
@functools.lru_cache(maxsize=None)
|
||||
def build_attention_mask(self):
|
||||
# lazily create causal attention mask, with full attention between the vision tokens
|
||||
# pytorch uses additive attention mask; fill with -inf
|
||||
mask = torch.empty(self.context_length, self.context_length)
|
||||
mask.fill_(float("-inf"))
|
||||
mask.triu_(1) # zero out the lower diagonal
|
||||
return mask
|
||||
|
||||
def tokenize(self, texts, context_length=77, truncate=True):
|
||||
"""
|
||||
Returns the tokenized representation of given input string(s)
|
||||
Parameters
|
||||
----------
|
||||
texts : Union[str, List[str]]
|
||||
An input string or a list of input strings to tokenize
|
||||
context_length : int
|
||||
The context length to use; all CLIP models use 77 as the context length
|
||||
truncate: bool
|
||||
Whether to truncate the text in case its encoding is longer than the context length
|
||||
Returns
|
||||
-------
|
||||
A two-dimensional tensor containing the resulting tokens, shape = [number of input strings, context_length].
|
||||
We return LongTensor when torch version is <1.8.0, since older index_select requires indices to be long.
|
||||
"""
|
||||
if isinstance(texts, str):
|
||||
texts = [texts]
|
||||
|
||||
sot_token = self._tokenizer.encoder["<|startoftext|>"]
|
||||
eot_token = self._tokenizer.encoder["<|endoftext|>"]
|
||||
all_tokens = [[sot_token] + self._tokenizer.encode(text) + [eot_token] for text in texts]
|
||||
if packaging.version.parse(torch.__version__) < packaging.version.parse("1.8.0"):
|
||||
result = torch.zeros(len(all_tokens), context_length, dtype=torch.long)
|
||||
else:
|
||||
result = torch.zeros(len(all_tokens), context_length, dtype=torch.int)
|
||||
|
||||
for i, tokens in enumerate(all_tokens):
|
||||
if len(tokens) > context_length:
|
||||
if truncate:
|
||||
tokens = tokens[:context_length]
|
||||
tokens[-1] = eot_token
|
||||
else:
|
||||
raise RuntimeError(f"Input {texts[i]} is too long for context length {context_length}")
|
||||
result[i, :len(tokens)] = torch.tensor(tokens)
|
||||
|
||||
return result
|
||||
|
||||
def forward(self, text):
|
||||
x = self.token_embedding(text) # [batch_size, n_ctx, d_model]
|
||||
|
||||
x = x + self.positional_embedding
|
||||
x = x.permute(1, 0, 2) # NLD -> LND
|
||||
x = self.transformer(x)
|
||||
x = x.permute(1, 0, 2) # LND -> NLD
|
||||
x = self.ln_final(x)
|
||||
|
||||
# x.shape = [batch_size, n_ctx, transformer.width]
|
||||
# take features from the eot embedding (eot_token is the highest number in each sequence)
|
||||
x = x[torch.arange(x.shape[0]), text.argmax(dim=-1)] @ self.text_projection
|
||||
|
||||
return x
|
||||
|
||||
|
||||
def clip_text_b16(
|
||||
embed_dim=512,
|
||||
context_length=77,
|
||||
vocab_size=49408,
|
||||
transformer_width=512,
|
||||
transformer_heads=8,
|
||||
transformer_layers=12,
|
||||
checkpoint_num=0,
|
||||
pretrained=True,
|
||||
):
|
||||
# raise NotImplementedError
|
||||
model = CLIP_TEXT(
|
||||
embed_dim,
|
||||
context_length,
|
||||
vocab_size,
|
||||
transformer_width,
|
||||
transformer_heads,
|
||||
transformer_layers,
|
||||
checkpoint_num,
|
||||
)
|
||||
# pretrained = _MODELS["ViT-B/16"]
|
||||
# logger.info(f"Load pretrained weights from {pretrained}")
|
||||
# state_dict = torch.load(pretrained, map_location='cpu')
|
||||
# model.load_state_dict(state_dict, strict=False)
|
||||
# return model.eval()
|
||||
if pretrained:
|
||||
if isinstance(pretrained, str) and pretrained != "bert-base-uncased":
|
||||
pretrained = _MODELS[pretrained]
|
||||
else:
|
||||
pretrained = _MODELS["ViT-B/16"]
|
||||
logger.info(f"Load pretrained weights from {pretrained}")
|
||||
state_dict = torch.load(pretrained, map_location='cpu')
|
||||
if context_length != state_dict["positional_embedding"].size(0):
|
||||
# assert context_length < state_dict["positional_embedding"].size(0), "Cannot increase context length."
|
||||
print(f"Resize positional embedding from {state_dict['positional_embedding'].size(0)} to {context_length}")
|
||||
if context_length < state_dict["positional_embedding"].size(0):
|
||||
state_dict["positional_embedding"] = state_dict["positional_embedding"][:context_length]
|
||||
else:
|
||||
state_dict["positional_embedding"] = F.pad(
|
||||
state_dict["positional_embedding"],
|
||||
(0, 0, 0, context_length - state_dict["positional_embedding"].size(0)),
|
||||
value=0,
|
||||
)
|
||||
|
||||
message = model.load_state_dict(state_dict, strict=False)
|
||||
print(f"Load pretrained weights from {pretrained}: {message}")
|
||||
return model.eval()
|
||||
|
||||
|
||||
def clip_text_l14(
|
||||
embed_dim=768,
|
||||
context_length=77,
|
||||
vocab_size=49408,
|
||||
transformer_width=768,
|
||||
transformer_heads=12,
|
||||
transformer_layers=12,
|
||||
checkpoint_num=0,
|
||||
pretrained=True,
|
||||
):
|
||||
model = CLIP_TEXT(
|
||||
embed_dim,
|
||||
context_length,
|
||||
vocab_size,
|
||||
transformer_width,
|
||||
transformer_heads,
|
||||
transformer_layers,
|
||||
checkpoint_num,
|
||||
)
|
||||
if pretrained:
|
||||
if isinstance(pretrained, str) and pretrained != "bert-base-uncased":
|
||||
pretrained = _MODELS[pretrained]
|
||||
else:
|
||||
pretrained = _MODELS["ViT-L/14"]
|
||||
logger.info(f"Load pretrained weights from {pretrained}")
|
||||
state_dict = torch.load(pretrained, map_location='cpu')
|
||||
if context_length != state_dict["positional_embedding"].size(0):
|
||||
# assert context_length < state_dict["positional_embedding"].size(0), "Cannot increase context length."
|
||||
print(f"Resize positional embedding from {state_dict['positional_embedding'].size(0)} to {context_length}")
|
||||
if context_length < state_dict["positional_embedding"].size(0):
|
||||
state_dict["positional_embedding"] = state_dict["positional_embedding"][:context_length]
|
||||
else:
|
||||
state_dict["positional_embedding"] = F.pad(
|
||||
state_dict["positional_embedding"],
|
||||
(0, 0, 0, context_length - state_dict["positional_embedding"].size(0)),
|
||||
value=0,
|
||||
)
|
||||
|
||||
message = model.load_state_dict(state_dict, strict=False)
|
||||
print(f"Load pretrained weights from {pretrained}: {message}")
|
||||
return model.eval()
|
||||
|
||||
|
||||
def clip_text_l14_336(
|
||||
embed_dim=768,
|
||||
context_length=77,
|
||||
vocab_size=49408,
|
||||
transformer_width=768,
|
||||
transformer_heads=12,
|
||||
transformer_layers=12,
|
||||
):
|
||||
raise NotImplementedError
|
||||
model = CLIP_TEXT(
|
||||
embed_dim,
|
||||
context_length,
|
||||
vocab_size,
|
||||
transformer_width,
|
||||
transformer_heads,
|
||||
transformer_layers
|
||||
)
|
||||
pretrained = _MODELS["ViT-L/14_336"]
|
||||
logger.info(f"Load pretrained weights from {pretrained}")
|
||||
state_dict = torch.load(pretrained, map_location='cpu')
|
||||
model.load_state_dict(state_dict, strict=False)
|
||||
return model.eval()
|
||||
|
||||
|
||||
def build_clip(config):
|
||||
model_cls = config.text_encoder.clip_teacher
|
||||
model = eval(model_cls)()
|
||||
return model
|
||||
@@ -0,0 +1,362 @@
|
||||
#!/usr/bin/env python
|
||||
import os
|
||||
import logging
|
||||
from collections import OrderedDict
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
from einops import rearrange
|
||||
from timm.models.layers import DropPath
|
||||
from timm.models.registry import register_model
|
||||
|
||||
import torch.utils.checkpoint as checkpoint
|
||||
|
||||
# from models.utils import load_temp_embed_with_mismatch
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
def load_temp_embed_with_mismatch(temp_embed_old, temp_embed_new, add_zero=True):
|
||||
"""
|
||||
Add/Remove extra temporal_embeddings as needed.
|
||||
https://arxiv.org/abs/2104.00650 shows adding zero paddings works.
|
||||
|
||||
temp_embed_old: (1, num_frames_old, 1, d)
|
||||
temp_embed_new: (1, num_frames_new, 1, d)
|
||||
add_zero: bool, if True, add zero, else, interpolate trained embeddings.
|
||||
"""
|
||||
# TODO zero pad
|
||||
num_frms_new = temp_embed_new.shape[1]
|
||||
num_frms_old = temp_embed_old.shape[1]
|
||||
logger.info(f"Load temporal_embeddings, lengths: {num_frms_old}-->{num_frms_new}")
|
||||
if num_frms_new > num_frms_old:
|
||||
if add_zero:
|
||||
temp_embed_new[
|
||||
:, :num_frms_old
|
||||
] = temp_embed_old # untrained embeddings are zeros.
|
||||
else:
|
||||
temp_embed_new = interpolate_temporal_pos_embed(temp_embed_old, num_frms_new)
|
||||
elif num_frms_new < num_frms_old:
|
||||
temp_embed_new = temp_embed_old[:, :num_frms_new]
|
||||
else: # =
|
||||
temp_embed_new = temp_embed_old
|
||||
return temp_embed_new
|
||||
|
||||
|
||||
# On P1, model extracted from https://huggingface.co/laion/CLIP-ViT-L-14-DataComp.XL-s13B-b90K
|
||||
MODEL_PATH = ''
|
||||
_MODELS = {
|
||||
"ViT-L/14": os.path.join(MODEL_PATH, "ViCLIP-L_InternVid-FLT-10M.pth"),
|
||||
"ViT-B/16": os.path.join(MODEL_PATH, "ViCLIP-B-InternVid-FLT-10M.pth"),
|
||||
}
|
||||
|
||||
|
||||
class QuickGELU(nn.Module):
|
||||
def forward(self, x):
|
||||
return x * torch.sigmoid(1.702 * x)
|
||||
|
||||
|
||||
class ResidualAttentionBlock(nn.Module):
|
||||
def __init__(self, d_model, n_head, drop_path=0., attn_mask=None, dropout=0.):
|
||||
super().__init__()
|
||||
|
||||
self.drop_path1 = DropPath(drop_path) if drop_path > 0. else nn.Identity()
|
||||
self.drop_path2 = DropPath(drop_path) if drop_path > 0. else nn.Identity()
|
||||
# logger.info(f'Droppath: {drop_path}')
|
||||
self.attn = nn.MultiheadAttention(d_model, n_head, dropout=dropout)
|
||||
self.ln_1 = nn.LayerNorm(d_model)
|
||||
self.mlp = nn.Sequential(OrderedDict([
|
||||
("c_fc", nn.Linear(d_model, d_model * 4)),
|
||||
("gelu", QuickGELU()),
|
||||
("drop1", nn.Dropout(dropout)),
|
||||
("c_proj", nn.Linear(d_model * 4, d_model)),
|
||||
("drop2", nn.Dropout(dropout)),
|
||||
]))
|
||||
self.ln_2 = nn.LayerNorm(d_model)
|
||||
self.attn_mask = attn_mask
|
||||
|
||||
def attention(self, x):
|
||||
self.attn_mask = self.attn_mask.to(dtype=x.dtype, device=x.device) if self.attn_mask is not None else None
|
||||
return self.attn(x, x, x, need_weights=False, attn_mask=self.attn_mask)[0]
|
||||
|
||||
def forward(self, x):
|
||||
x = x + self.drop_path1(self.attention(self.ln_1(x)))
|
||||
x = x + self.drop_path2(self.mlp(self.ln_2(x)))
|
||||
return x
|
||||
|
||||
|
||||
class Transformer(nn.Module):
|
||||
def __init__(self, width, layers, heads, drop_path=0., checkpoint_num=0, dropout=0.):
|
||||
super().__init__()
|
||||
dpr = [x.item() for x in torch.linspace(0, drop_path, layers)]
|
||||
self.resblocks = nn.ModuleList()
|
||||
for idx in range(layers):
|
||||
self.resblocks.append(ResidualAttentionBlock(width, heads, drop_path=dpr[idx], dropout=dropout))
|
||||
self.checkpoint_num = checkpoint_num
|
||||
|
||||
def forward(self, x):
|
||||
for idx, blk in enumerate(self.resblocks):
|
||||
if idx < self.checkpoint_num:
|
||||
x = checkpoint.checkpoint(blk, x)
|
||||
else:
|
||||
x = blk(x)
|
||||
return x
|
||||
|
||||
|
||||
class VisionTransformer(nn.Module):
|
||||
def __init__(
|
||||
self, input_resolution, patch_size, width, layers, heads, output_dim=None,
|
||||
kernel_size=1, num_frames=8, drop_path=0, checkpoint_num=0, dropout=0.,
|
||||
temp_embed=True,
|
||||
):
|
||||
super().__init__()
|
||||
self.output_dim = output_dim
|
||||
self.conv1 = nn.Conv3d(
|
||||
3, width,
|
||||
(kernel_size, patch_size, patch_size),
|
||||
(kernel_size, patch_size, patch_size),
|
||||
(0, 0, 0), bias=False
|
||||
)
|
||||
|
||||
scale = width ** -0.5
|
||||
self.class_embedding = nn.Parameter(scale * torch.randn(width))
|
||||
self.positional_embedding = nn.Parameter(scale * torch.randn((input_resolution // patch_size) ** 2 + 1, width))
|
||||
self.ln_pre = nn.LayerNorm(width)
|
||||
if temp_embed:
|
||||
self.temporal_positional_embedding = nn.Parameter(torch.zeros(1, num_frames, width))
|
||||
|
||||
self.transformer = Transformer(
|
||||
width, layers, heads, drop_path=drop_path, checkpoint_num=checkpoint_num,
|
||||
dropout=dropout)
|
||||
|
||||
self.ln_post = nn.LayerNorm(width)
|
||||
if output_dim is not None:
|
||||
self.proj = nn.Parameter(torch.empty(width, output_dim))
|
||||
else:
|
||||
self.proj = None
|
||||
|
||||
self.dropout = nn.Dropout(dropout)
|
||||
|
||||
def get_num_layers(self):
|
||||
return len(self.transformer.resblocks)
|
||||
|
||||
@torch.jit.ignore
|
||||
def no_weight_decay(self):
|
||||
return {'positional_embedding', 'class_embedding', 'temporal_positional_embedding'}
|
||||
|
||||
def mask_tokens(self, inputs, masking_prob=0.0):
|
||||
B, L, _ = inputs.shape
|
||||
|
||||
# This is different from text as we are masking a fix number of tokens
|
||||
Lm = int(masking_prob * L)
|
||||
masked_indices = torch.zeros(B, L)
|
||||
indices = torch.argsort(torch.rand_like(masked_indices), dim=-1)[:, :Lm]
|
||||
batch_indices = (
|
||||
torch.arange(masked_indices.shape[0]).unsqueeze(-1).expand_as(indices)
|
||||
)
|
||||
masked_indices[batch_indices, indices] = 1
|
||||
|
||||
masked_indices = masked_indices.bool()
|
||||
|
||||
return inputs[~masked_indices].reshape(B, -1, inputs.shape[-1])
|
||||
|
||||
def forward(self, x, masking_prob=0.0):
|
||||
x = self.conv1(x) # shape = [*, width, grid, grid]
|
||||
B, C, T, H, W = x.shape
|
||||
x = x.permute(0, 2, 3, 4, 1).reshape(B * T, H * W, C)
|
||||
|
||||
x = torch.cat([self.class_embedding.to(x.dtype) + torch.zeros(x.shape[0], 1, x.shape[-1], dtype=x.dtype, device=x.device), x], dim=1) # shape = [*, grid ** 2 + 1, width]
|
||||
x = x + self.positional_embedding.to(x.dtype)
|
||||
|
||||
# temporal pos
|
||||
cls_tokens = x[:B, :1, :]
|
||||
x = x[:, 1:]
|
||||
x = rearrange(x, '(b t) n m -> (b n) t m', b=B, t=T)
|
||||
if hasattr(self, 'temporal_positional_embedding'):
|
||||
if x.size(1) == 1:
|
||||
# This is a workaround for unused parameter issue
|
||||
x = x + self.temporal_positional_embedding.mean(1)
|
||||
else:
|
||||
x = x + self.temporal_positional_embedding
|
||||
x = rearrange(x, '(b n) t m -> b (n t) m', b=B, t=T)
|
||||
|
||||
if masking_prob > 0.0:
|
||||
x = self.mask_tokens(x, masking_prob)
|
||||
|
||||
x = torch.cat((cls_tokens, x), dim=1)
|
||||
|
||||
x = self.ln_pre(x)
|
||||
|
||||
x = x.permute(1, 0, 2) #BND -> NBD
|
||||
x = self.transformer(x)
|
||||
|
||||
x = self.ln_post(x)
|
||||
|
||||
if self.proj is not None:
|
||||
x = self.dropout(x[0]) @ self.proj
|
||||
else:
|
||||
x = x.permute(1, 0, 2) #NBD -> BND
|
||||
|
||||
return x
|
||||
|
||||
|
||||
def inflate_weight(weight_2d, time_dim, center=True):
|
||||
logger.info(f'Init center: {center}')
|
||||
if center:
|
||||
weight_3d = torch.zeros(*weight_2d.shape)
|
||||
weight_3d = weight_3d.unsqueeze(2).repeat(1, 1, time_dim, 1, 1)
|
||||
middle_idx = time_dim // 2
|
||||
weight_3d[:, :, middle_idx, :, :] = weight_2d
|
||||
else:
|
||||
weight_3d = weight_2d.unsqueeze(2).repeat(1, 1, time_dim, 1, 1)
|
||||
weight_3d = weight_3d / time_dim
|
||||
return weight_3d
|
||||
|
||||
|
||||
def load_state_dict(model, state_dict, input_resolution=224, patch_size=16, center=True):
|
||||
state_dict_3d = model.state_dict()
|
||||
for k in state_dict.keys():
|
||||
if k in state_dict_3d.keys() and state_dict[k].shape != state_dict_3d[k].shape:
|
||||
if len(state_dict_3d[k].shape) <= 2:
|
||||
logger.info(f'Ignore: {k}')
|
||||
continue
|
||||
logger.info(f'Inflate: {k}, {state_dict[k].shape} => {state_dict_3d[k].shape}')
|
||||
time_dim = state_dict_3d[k].shape[2]
|
||||
state_dict[k] = inflate_weight(state_dict[k], time_dim, center=center)
|
||||
|
||||
pos_embed_checkpoint = state_dict['positional_embedding']
|
||||
embedding_size = pos_embed_checkpoint.shape[-1]
|
||||
num_patches = (input_resolution // patch_size) ** 2
|
||||
orig_size = int((pos_embed_checkpoint.shape[-2] - 1) ** 0.5)
|
||||
new_size = int(num_patches ** 0.5)
|
||||
if orig_size != new_size:
|
||||
logger.info(f'Pos_emb from {orig_size} to {new_size}')
|
||||
extra_tokens = pos_embed_checkpoint[:1]
|
||||
pos_tokens = pos_embed_checkpoint[1:]
|
||||
pos_tokens = pos_tokens.reshape(-1, orig_size, orig_size, embedding_size).permute(0, 3, 1, 2)
|
||||
pos_tokens = torch.nn.functional.interpolate(
|
||||
pos_tokens, size=(new_size, new_size), mode='bicubic', align_corners=False)
|
||||
pos_tokens = pos_tokens.permute(0, 2, 3, 1).flatten(0, 2)
|
||||
new_pos_embed = torch.cat((extra_tokens, pos_tokens), dim=0)
|
||||
state_dict['positional_embedding'] = new_pos_embed
|
||||
|
||||
message = model.load_state_dict(state_dict, strict=False)
|
||||
logger.info(f"Load pretrained weights: {message}")
|
||||
|
||||
|
||||
@register_model
|
||||
def clip_joint_b16(
|
||||
pretrained=False, input_resolution=224, kernel_size=1,
|
||||
center=True, num_frames=8, drop_path=0., checkpoint_num=0,
|
||||
dropout=0.,
|
||||
):
|
||||
model = VisionTransformer(
|
||||
input_resolution=input_resolution, patch_size=16,
|
||||
width=768, layers=12, heads=12, output_dim=512,
|
||||
kernel_size=kernel_size, num_frames=num_frames,
|
||||
drop_path=drop_path, checkpoint_num=checkpoint_num,
|
||||
dropout=dropout,
|
||||
)
|
||||
# raise NotImplementedError
|
||||
if pretrained:
|
||||
if isinstance(pretrained, str):
|
||||
model_name = pretrained
|
||||
else:
|
||||
model_name = "ViT-B/16"
|
||||
|
||||
logger.info('load pretrained weights')
|
||||
state_dict = torch.load(_MODELS[model_name], map_location='cpu')
|
||||
load_state_dict(model, state_dict, input_resolution=input_resolution, patch_size=16, center=center)
|
||||
return model.eval()
|
||||
|
||||
|
||||
@register_model
|
||||
def clip_joint_l14(
|
||||
pretrained=False, input_resolution=224, kernel_size=1,
|
||||
center=True, num_frames=8, drop_path=0., checkpoint_num=0,
|
||||
dropout=0.,
|
||||
):
|
||||
model = VisionTransformer(
|
||||
input_resolution=input_resolution, patch_size=14,
|
||||
width=1024, layers=24, heads=16, output_dim=768,
|
||||
kernel_size=kernel_size, num_frames=num_frames,
|
||||
drop_path=drop_path, checkpoint_num=checkpoint_num,
|
||||
dropout=dropout,
|
||||
)
|
||||
|
||||
if pretrained:
|
||||
if isinstance(pretrained, str):
|
||||
model_name = pretrained
|
||||
else:
|
||||
model_name = "ViT-L/14"
|
||||
logger.info('load pretrained weights')
|
||||
state_dict = torch.load(_MODELS[model_name], map_location='cpu')
|
||||
load_state_dict(model, state_dict, input_resolution=input_resolution, patch_size=14, center=center)
|
||||
return model.eval()
|
||||
|
||||
|
||||
@register_model
|
||||
def clip_joint_l14_336(
|
||||
pretrained=True, input_resolution=336, kernel_size=1,
|
||||
center=True, num_frames=8, drop_path=0.
|
||||
):
|
||||
raise NotImplementedError
|
||||
model = VisionTransformer(
|
||||
input_resolution=input_resolution, patch_size=14,
|
||||
width=1024, layers=24, heads=16, output_dim=768,
|
||||
kernel_size=kernel_size, num_frames=num_frames,
|
||||
drop_path=drop_path,
|
||||
)
|
||||
if pretrained:
|
||||
logger.info('load pretrained weights')
|
||||
state_dict = torch.load(_MODELS["ViT-L/14_336"], map_location='cpu')
|
||||
load_state_dict(model, state_dict, input_resolution=input_resolution, patch_size=14, center=center)
|
||||
return model.eval()
|
||||
|
||||
|
||||
def interpolate_pos_embed_vit(state_dict, new_model):
|
||||
key = "vision_encoder.temporal_positional_embedding"
|
||||
if key in state_dict:
|
||||
vision_temp_embed_new = new_model.state_dict()[key]
|
||||
vision_temp_embed_new = vision_temp_embed_new.unsqueeze(2) # [1, n, d] -> [1, n, 1, d]
|
||||
vision_temp_embed_old = state_dict[key]
|
||||
vision_temp_embed_old = vision_temp_embed_old.unsqueeze(2)
|
||||
|
||||
state_dict[key] = load_temp_embed_with_mismatch(
|
||||
vision_temp_embed_old, vision_temp_embed_new, add_zero=False
|
||||
).squeeze(2)
|
||||
|
||||
key = "text_encoder.positional_embedding"
|
||||
if key in state_dict:
|
||||
text_temp_embed_new = new_model.state_dict()[key]
|
||||
text_temp_embed_new = text_temp_embed_new.unsqueeze(0).unsqueeze(2) # [n, d] -> [1, n, 1, d]
|
||||
text_temp_embed_old = state_dict[key]
|
||||
text_temp_embed_old = text_temp_embed_old.unsqueeze(0).unsqueeze(2)
|
||||
|
||||
state_dict[key] = load_temp_embed_with_mismatch(
|
||||
text_temp_embed_old, text_temp_embed_new, add_zero=False
|
||||
).squeeze(2).squeeze(0)
|
||||
return state_dict
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
import time
|
||||
from fvcore.nn import FlopCountAnalysis
|
||||
from fvcore.nn import flop_count_table
|
||||
import numpy as np
|
||||
|
||||
seed = 4217
|
||||
np.random.seed(seed)
|
||||
torch.manual_seed(seed)
|
||||
torch.cuda.manual_seed(seed)
|
||||
torch.cuda.manual_seed_all(seed)
|
||||
num_frames = 8
|
||||
|
||||
# model = clip_joint_b16(pretrained=True, kernel_size=1, num_frames=8, num_classes=400, drop_path=0.1)
|
||||
# logger.info(model)
|
||||
model = clip_joint_l14(pretrained=False)
|
||||
|
||||
flops = FlopCountAnalysis(model, torch.rand(1, 3, num_frames, 224, 224))
|
||||
s = time.time()
|
||||
logger.info(flop_count_table(flops, max_depth=1))
|
||||
logger.info(time.time()-s)
|
||||
# logger.info(model(torch.rand(1, 3, num_frames, 224, 224)).shape)
|
||||
@@ -0,0 +1,101 @@
|
||||
import os
|
||||
from typing import Optional
|
||||
from pathlib import Path
|
||||
|
||||
from func_timeout import func_timeout, FunctionTimedOut
|
||||
from PIL import Image
|
||||
from torch.utils.data import Dataset, DataLoader
|
||||
|
||||
from .logger import logger
|
||||
from .video_utils import extract_frames
|
||||
|
||||
|
||||
ALL_VIDEO_EXT = set(["mp4", "webm", "mkv", "avi", "flv", "mov"])
|
||||
VIDEO_READER_TIMEOUT = 300
|
||||
|
||||
|
||||
def collate_fn(batch):
|
||||
batch = list(filter(lambda x: x is not None, batch))
|
||||
if len(batch) != 0:
|
||||
return {k: [item[k] for item in batch] for k in batch[0].keys()}
|
||||
return {}
|
||||
|
||||
|
||||
class VideoDataset(Dataset):
|
||||
def __init__(
|
||||
self,
|
||||
dataset_inputs: dict[str, list[str]],
|
||||
video_folder: Optional[str] = None,
|
||||
video_path_column: str = "video_path",
|
||||
text_column: Optional[str] = None,
|
||||
sample_method: str = "mid",
|
||||
num_sampled_frames: int = 1,
|
||||
num_sample_stride: Optional[int] = None
|
||||
):
|
||||
length = len(dataset_inputs[list(dataset_inputs.keys())[0]])
|
||||
if not all(len(v) == length for v in dataset_inputs.values()):
|
||||
raise ValueError("All values in the dataset_inputs must have the same length.")
|
||||
|
||||
self.video_path_column = video_path_column
|
||||
self.video_folder = video_folder
|
||||
self.video_path_list = dataset_inputs[video_path_column]
|
||||
if self.video_folder is not None:
|
||||
self.video_path_list = [os.path.join(self.video_folder, video_path) for video_path in self.video_path_list]
|
||||
self.text_column = text_column
|
||||
self.text_list = dataset_inputs[self.text_column] if self.text_column is not None else None
|
||||
|
||||
self.sample_method = sample_method
|
||||
self.num_sampled_frames = num_sampled_frames
|
||||
self.num_sample_stride = num_sample_stride
|
||||
|
||||
def __getitem__(self, index):
|
||||
video_path = self.video_path_list[index]
|
||||
if self.sample_method == "image":
|
||||
try:
|
||||
sampled_frame_idx_list = None
|
||||
with open(video_path, "rb") as f:
|
||||
sampled_frame_list = [Image.open(f).convert("RGB")]
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to extract frames from video {video_path}. Error is {e}.")
|
||||
return None
|
||||
else:
|
||||
# It is a trick to deal with decord hanging when reading some abnormal videos.
|
||||
try:
|
||||
sample_args = (video_path, self.sample_method, self.num_sampled_frames, self.num_sample_stride)
|
||||
sampled_frame_idx_list, sampled_frame_list = func_timeout(
|
||||
VIDEO_READER_TIMEOUT, extract_frames, args=sample_args
|
||||
)
|
||||
except FunctionTimedOut:
|
||||
logger.warning(f"Read {video_path} timeout.")
|
||||
return None
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to extract frames from video {video_path}. Error is {e}.")
|
||||
return None
|
||||
|
||||
item = {
|
||||
"path": video_path,
|
||||
"sampled_frame_idx": sampled_frame_idx_list,
|
||||
"sampled_frame": sampled_frame_list,
|
||||
}
|
||||
if self.text_list is not None:
|
||||
item["text"] = self.text_list[index]
|
||||
|
||||
return item
|
||||
|
||||
def __len__(self):
|
||||
return len(self.video_path_list)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
video_folder = Path("your_video_folder")
|
||||
video_path_list = []
|
||||
for ext in ALL_VIDEO_EXT:
|
||||
video_path_list += [str(file.relative_to(video_folder)) for file in video_folder.glob(f"*.{ext}")]
|
||||
|
||||
video_dataset = VideoDataset(dataset_inputs={"video_path": video_path_list})
|
||||
video_dataloader = DataLoader(
|
||||
video_dataset, batch_size=16, num_workers=16, collate_fn=collate_fn
|
||||
)
|
||||
for idx, batch in enumerate(video_dataloader):
|
||||
if len(batch) != 0:
|
||||
print(batch["video_path"], batch["sampled_frame_idx"], len(batch["video_path"]))
|
||||
@@ -0,0 +1,120 @@
|
||||
import os
|
||||
from typing import List
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from PIL import Image
|
||||
from torchvision.datasets.utils import download_url
|
||||
|
||||
from .longclip import longclip
|
||||
from .viclip import get_viclip
|
||||
from .video_utils import extract_frames
|
||||
|
||||
# All metrics.
|
||||
__all__ = ["VideoCLIPXLScore"]
|
||||
|
||||
_MODELS = {
|
||||
"ViClip-InternVid-10M-FLT": "https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/video_caption/clip/ViClip-InternVid-10M-FLT.pth",
|
||||
"LongCLIP-L": "https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/video_caption/clip/longclip-L.pt",
|
||||
"VideoCLIP-XL-v2": "https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/video_caption/clip/VideoCLIP-XL-v2.bin",
|
||||
}
|
||||
_MD5 = {
|
||||
"ViClip-InternVid-10M-FLT": "b1ebf538225438b3b75e477da7735cd0",
|
||||
"LongCLIP-L": "5478b662f6f85ca0ebd4bb05f9b592f3",
|
||||
"VideoCLIP-XL-v2": "cebda0bab14b677ec061a57e80791f35",
|
||||
}
|
||||
|
||||
def normalize(
|
||||
data: np.array,
|
||||
mean: list[float] = [0.485, 0.456, 0.406],
|
||||
std: list[float] = [0.229, 0.224, 0.225]
|
||||
):
|
||||
v_mean = np.array(mean).reshape(1, 1, 3)
|
||||
v_std = np.array(std).reshape(1, 1, 3)
|
||||
|
||||
return (data / 255.0 - v_mean) / v_std
|
||||
|
||||
|
||||
class VideoCLIPXL(nn.Module):
|
||||
def __init__(self, root: str = "~/.cache/clip"):
|
||||
super(VideoCLIPXL, self).__init__()
|
||||
|
||||
self.root = os.path.expanduser(root)
|
||||
if not os.path.exists(self.root):
|
||||
os.makedirs(self.root)
|
||||
|
||||
k = "LongCLIP-L"
|
||||
filename = os.path.basename(_MODELS[k])
|
||||
download_url(_MODELS[k], self.root, filename=filename, md5=_MD5[k])
|
||||
self.model = longclip.load(os.path.join(self.root, filename), device="cpu")[0].float()
|
||||
|
||||
k = "ViClip-InternVid-10M-FLT"
|
||||
filename = os.path.basename(_MODELS[k])
|
||||
download_url(_MODELS[k], self.root, filename=filename, md5=_MD5[k])
|
||||
self.viclip_model = get_viclip("l", os.path.join(self.root, filename))["viclip"].float()
|
||||
|
||||
# delete unused encoder
|
||||
del self.model.visual
|
||||
del self.viclip_model.text_encoder
|
||||
|
||||
|
||||
class VideoCLIPXLScore():
|
||||
def __init__(self, root: str = "~/.cache/clip", device: str = "cpu"):
|
||||
self.root = os.path.expanduser(root)
|
||||
if not os.path.exists(self.root):
|
||||
os.makedirs(self.root)
|
||||
|
||||
k = "VideoCLIP-XL-v2"
|
||||
filename = os.path.basename(_MODELS[k])
|
||||
download_url(_MODELS[k], self.root, filename=filename, md5=_MD5[k])
|
||||
self.model = VideoCLIPXL()
|
||||
state_dict = torch.load(os.path.join(self.root, filename), map_location="cpu")
|
||||
self.model.load_state_dict(state_dict)
|
||||
self.model.to(device)
|
||||
|
||||
self.device = device
|
||||
|
||||
def __call__(self, videos: List[List[Image.Image]], texts: List[str]):
|
||||
assert len(videos) == len(texts)
|
||||
|
||||
# Use cv2.resize in accordance with the official demo. Resize and Normalize => B * [T, 224, 224, 3].
|
||||
videos = [[cv2.cvtColor(np.array(f), cv2.COLOR_RGB2BGR) for f in v] for v in videos]
|
||||
resize_videos = [[cv2.resize(f, (224, 224)) for f in v] for v in videos]
|
||||
resize_normalizied_videos = [normalize(np.stack(v)) for v in resize_videos]
|
||||
|
||||
video_inputs = torch.stack([torch.from_numpy(v) for v in resize_normalizied_videos])
|
||||
video_inputs = video_inputs.float().permute(0, 1, 4, 2, 3).to(self.device, non_blocking=True) # BTCHW
|
||||
|
||||
with torch.no_grad():
|
||||
vid_features = torch.stack(
|
||||
[self.model.viclip_model.get_vid_features(x.unsqueeze(0)).float() for x in video_inputs]
|
||||
)
|
||||
vid_features.squeeze_()
|
||||
# vid_features = self.model.viclip_model.get_vid_features(video_inputs).float()
|
||||
text_inputs = longclip.tokenize(texts, truncate=True).to(self.device)
|
||||
text_features = self.model.model.encode_text(text_inputs)
|
||||
text_features = text_features / text_features.norm(dim=1, keepdim=True)
|
||||
scores = text_features @ vid_features.T
|
||||
|
||||
return scores.tolist() if len(videos) == 1 else scores.diagonal().tolist()
|
||||
|
||||
def __repr__(self):
|
||||
return "videoclipxl_score"
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
videos = ["your_video_path"] * 3
|
||||
texts = [
|
||||
"a joker",
|
||||
"glasses and flower",
|
||||
"The video opens with a view of a white building with multiple windows, partially obscured by leafless tree branches. The scene transitions to a closer view of the same building, with the tree branches more prominent in the foreground. The focus then shifts to a street sign that reads 'Abesses' in bold, yellow letters against a green background. The sign is attached to a metal structure, possibly a tram or bus stop. The sign is illuminated by a light source above it, and the background reveals a glimpse of the building and tree branches from earlier shots. The colors are muted, with the yellow sign standing out against the grey and green hues."
|
||||
]
|
||||
|
||||
video_clip_xl_score = VideoCLIPXLScore(device="cuda")
|
||||
batch_frames = []
|
||||
for v in videos:
|
||||
sampled_frames = extract_frames(v, sample_method="uniform", num_sampled_frames=8)[1]
|
||||
batch_frames.append(sampled_frames)
|
||||
print(video_clip_xl_score(batch_frames, texts))
|
||||
@@ -0,0 +1,44 @@
|
||||
import gc
|
||||
import random
|
||||
from contextlib import contextmanager
|
||||
from typing import List, Tuple, Optional
|
||||
|
||||
import numpy as np
|
||||
from decord import VideoReader
|
||||
from PIL import Image
|
||||
|
||||
|
||||
@contextmanager
|
||||
def video_reader(*args, **kwargs):
|
||||
"""A context manager to solve the memory leak of decord.
|
||||
"""
|
||||
vr = VideoReader(*args, **kwargs)
|
||||
try:
|
||||
yield vr
|
||||
finally:
|
||||
del vr
|
||||
gc.collect()
|
||||
|
||||
|
||||
def extract_frames(
|
||||
video_path: str,
|
||||
sample_method: str = "mid",
|
||||
num_sampled_frames: int = -1,
|
||||
sample_stride: int = -1,
|
||||
**kwargs
|
||||
) -> Optional[Tuple[List[int], List[Image.Image]]]:
|
||||
with video_reader(video_path, num_threads=2, **kwargs) as vr:
|
||||
if sample_method == "mid":
|
||||
sampled_frame_idx_list = [len(vr) // 2]
|
||||
elif sample_method == "uniform":
|
||||
sampled_frame_idx_list = np.linspace(0, len(vr), num_sampled_frames, endpoint=False, dtype=int)
|
||||
elif sample_method == "random":
|
||||
clip_length = min(len(vr), (num_sampled_frames - 1) * sample_stride + 1)
|
||||
start_idx = random.randint(0, len(vr) - clip_length)
|
||||
sampled_frame_idx_list = np.linspace(start_idx, start_idx + clip_length - 1, num_sampled_frames, dtype=int)
|
||||
else:
|
||||
raise ValueError(f"The sample_method {sample_method} must be mid, uniform or random.")
|
||||
sampled_frame_list = vr.get_batch(sampled_frame_idx_list).asnumpy()
|
||||
sampled_frame_list = [Image.fromarray(frame) for frame in sampled_frame_list]
|
||||
|
||||
return list(sampled_frame_idx_list), sampled_frame_list
|
||||
@@ -0,0 +1,169 @@
|
||||
import argparse
|
||||
import os
|
||||
import subprocess
|
||||
from datetime import datetime, timedelta
|
||||
from pathlib import Path
|
||||
from multiprocessing import Pool
|
||||
|
||||
import pandas as pd
|
||||
from tqdm import tqdm
|
||||
|
||||
from utils.logger import logger
|
||||
|
||||
|
||||
MIN_SECONDS = int(os.getenv("MIN_SECONDS", 3))
|
||||
MAX_SECONDS = int(os.getenv("MAX_SECONDS", 10))
|
||||
|
||||
|
||||
def get_command(start_time, video_path, video_duration, output_path):
|
||||
# Use FFmpeg to split the video. Re-encoding is needed to ensure the accuracy of the clip
|
||||
# at the cost of consuming computational resources.
|
||||
return [
|
||||
'ffmpeg',
|
||||
'-hide_banner',
|
||||
'-loglevel', 'panic',
|
||||
'-ss', str(start_time.time()),
|
||||
'-i', video_path,
|
||||
'-t', str(video_duration),
|
||||
'-c:v', 'libx264',
|
||||
'-preset', 'veryfast',
|
||||
'-crf', '22',
|
||||
'-c:a', 'aac',
|
||||
'-sn',
|
||||
output_path
|
||||
]
|
||||
|
||||
|
||||
def clip_video_star(args):
|
||||
return clip_video(*args)
|
||||
|
||||
|
||||
def clip_video(video_path, timecode_list, output_folder, video_duration):
|
||||
"""Recursively clip the video within the range of [MIN_SECONDS, MAX_SECONDS],
|
||||
according to the timecode obtained from cogvideox/video_caption/cutscene_detect.py.
|
||||
"""
|
||||
try:
|
||||
video_name = Path(video_path).stem
|
||||
|
||||
if len(timecode_list) == 0: # The video of a single scene.
|
||||
splitted_timecode_list = []
|
||||
start_time = datetime.strptime("00:00:00.000", "%H:%M:%S.%f")
|
||||
end_time = datetime.strptime(video_duration, "%H:%M:%S.%f")
|
||||
cur_start = start_time
|
||||
splitted_index = 0
|
||||
while cur_start < end_time:
|
||||
cur_end = min(cur_start + timedelta(seconds=MAX_SECONDS), end_time)
|
||||
cur_video_duration = (cur_end - cur_start).total_seconds()
|
||||
if cur_video_duration < MIN_SECONDS:
|
||||
cur_start = cur_end
|
||||
splitted_index += 1
|
||||
continue
|
||||
splitted_timecode_list.append([cur_start.strftime("%H:%M:%S.%f")[:-3], cur_end.strftime("%H:%M:%S.%f")[:-3]])
|
||||
output_path = os.path.join(output_folder, video_name + f"_{splitted_index}.mp4")
|
||||
if os.path.exists(output_path):
|
||||
logger.info(f"The clipped video {output_path} exists.")
|
||||
cur_start = cur_end
|
||||
splitted_index += 1
|
||||
continue
|
||||
else:
|
||||
command = get_command(cur_start, video_path, cur_video_duration, output_path)
|
||||
try:
|
||||
subprocess.run(command, check=True)
|
||||
except Exception as e:
|
||||
logger.warning(f"Run {command} error: {e}.")
|
||||
finally:
|
||||
cur_start = cur_end
|
||||
splitted_index += 1
|
||||
|
||||
for i, timecode in enumerate(timecode_list): # The video of multiple scenes.
|
||||
start_time = datetime.strptime(timecode[0], "%H:%M:%S.%f")
|
||||
end_time = datetime.strptime(timecode[1], "%H:%M:%S.%f")
|
||||
video_duration = (end_time - start_time).total_seconds()
|
||||
output_path = os.path.join(output_folder, video_name + f"_{i}.mp4")
|
||||
if os.path.exists(output_path):
|
||||
logger.info(f"The clipped video {output_path} exists.")
|
||||
continue
|
||||
if video_duration < MIN_SECONDS:
|
||||
continue
|
||||
if video_duration > MAX_SECONDS:
|
||||
splitted_timecode_list = []
|
||||
cur_start = start_time
|
||||
splitted_index = 0
|
||||
while cur_start < end_time:
|
||||
cur_end = min(cur_start + timedelta(seconds=MAX_SECONDS), end_time)
|
||||
cur_video_duration = (cur_end - cur_start).total_seconds()
|
||||
if cur_video_duration < MIN_SECONDS:
|
||||
break
|
||||
splitted_timecode_list.append([cur_start.strftime("%H:%M:%S.%f")[:-3], cur_end.strftime("%H:%M:%S.%f")[:-3]])
|
||||
splitted_output_path = os.path.join(output_folder, video_name + f"_{i}_{splitted_index}.mp4")
|
||||
if os.path.exists(splitted_output_path):
|
||||
logger.info(f"The clipped video {splitted_output_path} exists.")
|
||||
cur_start = cur_end
|
||||
splitted_index += 1
|
||||
continue
|
||||
else:
|
||||
command = get_command(cur_start, video_path, cur_video_duration, splitted_output_path)
|
||||
try:
|
||||
subprocess.run(command, check=True)
|
||||
except Exception as e:
|
||||
logger.warning(f"Run {command} error: {e}.")
|
||||
finally:
|
||||
cur_start = cur_end
|
||||
splitted_index += 1
|
||||
|
||||
continue
|
||||
|
||||
# We found that the current scene detected by PySceneDetect includes a few frames from
|
||||
# the next scene occasionally. Directly discard the last few frames of the current scene.
|
||||
video_duration = video_duration - 0.5
|
||||
command = get_command(start_time, video_path, video_duration, output_path)
|
||||
subprocess.run(command, check=True)
|
||||
except Exception as e:
|
||||
logger.warning(f"Clip video with {video_path}. Error is: {e}.")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser(description="Video Splitting")
|
||||
parser.add_argument(
|
||||
"--video_metadata_path", type=str, default=None, help="The path to the video dataset metadata (csv/jsonl)."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--video_path_column",
|
||||
type=str,
|
||||
default="video_path",
|
||||
help="The column contains the video path (an absolute path or a relative path w.r.t the video_folder).",
|
||||
)
|
||||
parser.add_argument("--video_folder", type=str, default="", help="The video folder.")
|
||||
parser.add_argument("--output_folder", type=str, default="outputs")
|
||||
parser.add_argument("--n_jobs", type=int, default=16)
|
||||
|
||||
parser.add_argument("--resolution_threshold", type=float, default=0, help="The resolution threshold.")
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
video_metadata_df = pd.read_json(args.video_metadata_path, lines=True)
|
||||
num_videos = len(video_metadata_df)
|
||||
video_metadata_df["resolution"] = video_metadata_df["frame_size"].apply(lambda x: x[0] * x[1])
|
||||
video_metadata_df = video_metadata_df[video_metadata_df["resolution"] >= args.resolution_threshold]
|
||||
logger.info(f"Filter {num_videos - len(video_metadata_df)} videos with resolution smaller than {args.resolution_threshold}.")
|
||||
video_path_list = video_metadata_df[args.video_path_column].to_list()
|
||||
video_id_list = [Path(video_path).stem for video_path in video_path_list]
|
||||
if len(video_id_list) != len(list(set(video_id_list))):
|
||||
logger.warning("Duplicate file names exist in the input video path list.")
|
||||
video_path_list = [os.path.join(args.video_folder, video_path) for video_path in video_path_list]
|
||||
video_timecode_list = video_metadata_df["timecode_list"].to_list()
|
||||
video_duration_list = video_metadata_df["duration"].to_list()
|
||||
|
||||
assert len(video_path_list) == len(video_timecode_list)
|
||||
os.makedirs(args.output_folder, exist_ok=True)
|
||||
args_list = [
|
||||
(video_path, timecode_list, args.output_folder, video_duration)
|
||||
for video_path, timecode_list, video_duration in zip(
|
||||
video_path_list, video_timecode_list, video_duration_list
|
||||
)
|
||||
]
|
||||
with Pool(args.n_jobs) as pool:
|
||||
# results = list(tqdm(pool.imap(clip_video_star, args_list), total=len(video_path_list)))
|
||||
results = pool.imap(clip_video_star, args_list)
|
||||
for result in tqdm(results, total=len(video_path_list)):
|
||||
pass
|
||||
@@ -0,0 +1,354 @@
|
||||
# Modified from https://github.com/mit-han-lab/llm-awq/blob/main/tinychat/vlm_demo_new.py.
|
||||
import argparse
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
import torch
|
||||
from accelerate import load_checkpoint_and_dispatch, PartialState
|
||||
from accelerate.utils import gather_object
|
||||
from decord import VideoReader
|
||||
from PIL import Image
|
||||
from natsort import natsorted
|
||||
from tqdm import tqdm
|
||||
from transformers import AutoConfig, AutoTokenizer
|
||||
|
||||
import tinychat.utils.constants
|
||||
# from tinychat.models.llava_llama import LlavaLlamaForCausalLM
|
||||
from tinychat.models.vila_llama import VilaLlamaForCausalLM
|
||||
from tinychat.stream_generators.llava_stream_gen import LlavaStreamGenerator
|
||||
from tinychat.utils.conversation_utils import gen_params
|
||||
from tinychat.utils.llava_image_processing import process_images
|
||||
from tinychat.utils.prompt_templates import (
|
||||
get_image_token,
|
||||
get_prompter,
|
||||
get_stop_token_ids,
|
||||
)
|
||||
from tinychat.utils.tune import (
|
||||
device_warmup,
|
||||
tune_llava_patch_embedding,
|
||||
)
|
||||
|
||||
from utils.filter import filter
|
||||
from utils.logger import logger
|
||||
|
||||
gen_params.seed = 1
|
||||
gen_params.temp = 1.0
|
||||
gen_params.top_p = 1.0
|
||||
|
||||
|
||||
def extract_uniform_frames(video_path: str, num_sampled_frames: int = 8):
|
||||
vr = VideoReader(video_path)
|
||||
sampled_frame_idx_list = np.linspace(0, len(vr), num_sampled_frames, endpoint=False, dtype=int)
|
||||
sampled_frame_list = []
|
||||
for idx in sampled_frame_idx_list:
|
||||
sampled_frame = Image.fromarray(vr[idx].asnumpy())
|
||||
sampled_frame_list.append(sampled_frame)
|
||||
|
||||
return sampled_frame_list
|
||||
|
||||
|
||||
def stream_output(output_stream):
|
||||
for outputs in output_stream:
|
||||
output_text = outputs["text"]
|
||||
output_text = output_text.strip().split(" ")
|
||||
# print(f"output_text: {output_text}.")
|
||||
return " ".join(output_text)
|
||||
|
||||
|
||||
def skip(*args, **kwargs):
|
||||
pass
|
||||
|
||||
|
||||
def parse_args():
|
||||
parser = argparse.ArgumentParser(description="Recaption videos with VILA1.5.")
|
||||
parser.add_argument(
|
||||
"--video_metadata_path",
|
||||
type=str,
|
||||
default=None,
|
||||
help="The path to the video dataset metadata (csv/jsonl).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--video_path_column",
|
||||
type=str,
|
||||
default="video_path",
|
||||
help="The column contains the video path (an absolute path or a relative path w.r.t the video_folder).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--caption_column",
|
||||
type=str,
|
||||
default="caption",
|
||||
help="The column contains the caption.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--video_folder", type=str, default="", help="The video folder."
|
||||
)
|
||||
parser.add_argument("--input_prompt", type=str, default="<video>\\n Elaborate on the visual and narrative elements of the video in detail.")
|
||||
parser.add_argument(
|
||||
"--model_type", type=str, default="LLaMa", help="type of the model"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--model_path", type=str, default="Efficient-Large-Model/Llama-3-VILA1.5-8b-AWQ"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--quant_path",
|
||||
type=str,
|
||||
default=None,
|
||||
)
|
||||
parser.add_argument(
|
||||
"--precision", type=str, default="W4A16", help="compute precision"
|
||||
)
|
||||
parser.add_argument("--num_sampled_frames", type=int, default=8)
|
||||
parser.add_argument(
|
||||
"--saved_path",
|
||||
type=str,
|
||||
required=True,
|
||||
help="The save path to the output results (csv/jsonl).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--saved_freq",
|
||||
type=int,
|
||||
default=100,
|
||||
help="The frequency to save the output results.",
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
"--basic_metadata_path", type=str, default=None, help="The path to the basic metadata (csv/jsonl)."
|
||||
)
|
||||
parser.add_argument("--min_resolution", type=float, default=0, help="The resolution threshold.")
|
||||
parser.add_argument("--min_duration", type=float, default=-1, help="The minimum duration.")
|
||||
parser.add_argument("--max_duration", type=float, default=-1, help="The maximum duration.")
|
||||
parser.add_argument(
|
||||
"--asethetic_score_metadata_path", type=str, default=None, help="The path to the video quality metadata (csv/jsonl)."
|
||||
)
|
||||
parser.add_argument("--min_asethetic_score", type=float, default=4.0, help="The asethetic score threshold.")
|
||||
parser.add_argument(
|
||||
"--asethetic_score_siglip_metadata_path", type=str, default=None, help="The path to the video quality metadata (csv/jsonl)."
|
||||
)
|
||||
parser.add_argument("--min_asethetic_score_siglip", type=float, default=4.0, help="The asethetic score (SigLIP) threshold.")
|
||||
parser.add_argument(
|
||||
"--text_score_metadata_path", type=str, default=None, help="The path to the video text score metadata (csv/jsonl)."
|
||||
)
|
||||
parser.add_argument("--min_text_score", type=float, default=0.02, help="The text threshold.")
|
||||
parser.add_argument(
|
||||
"--motion_score_metadata_path", type=str, default=None, help="The path to the video motion score metadata (csv/jsonl)."
|
||||
)
|
||||
parser.add_argument("--min_motion_score", type=float, default=2, help="The motion threshold.")
|
||||
|
||||
args = parser.parse_args()
|
||||
return args
|
||||
|
||||
|
||||
def main(args):
|
||||
if args.video_metadata_path.endswith(".csv"):
|
||||
video_metadata_df = pd.read_csv(args.video_metadata_path)
|
||||
elif args.video_metadata_path.endswith(".jsonl"):
|
||||
video_metadata_df = pd.read_json(args.video_metadata_path, lines=True)
|
||||
else:
|
||||
raise ValueError("The video_metadata_path must end with .csv or .jsonl.")
|
||||
video_path_list = video_metadata_df[args.video_path_column].tolist()
|
||||
video_path_list = [os.path.basename(video_path) for video_path in video_path_list]
|
||||
|
||||
if not (args.saved_path.endswith(".csv") or args.saved_path.endswith(".jsonl")):
|
||||
raise ValueError("The saved_path must end with .csv or .jsonl.")
|
||||
|
||||
if os.path.exists(args.saved_path):
|
||||
if args.saved_path.endswith(".csv"):
|
||||
saved_metadata_df = pd.read_csv(args.saved_path)
|
||||
elif args.saved_path.endswith(".jsonl"):
|
||||
saved_metadata_df = pd.read_json(args.saved_path, lines=True)
|
||||
saved_video_path_list = saved_metadata_df[args.video_path_column].tolist()
|
||||
video_path_list = list(set(video_path_list).difference(set(saved_video_path_list)))
|
||||
logger.info(
|
||||
f"Resume from {args.saved_path}: {len(saved_video_path_list)} processed and {len(video_path_list)} to be processed."
|
||||
)
|
||||
|
||||
video_path_list = filter(
|
||||
video_path_list,
|
||||
basic_metadata_path=args.basic_metadata_path,
|
||||
min_resolution=args.min_resolution,
|
||||
min_duration=args.min_duration,
|
||||
max_duration=args.max_duration,
|
||||
asethetic_score_metadata_path=args.asethetic_score_metadata_path,
|
||||
min_asethetic_score=args.min_asethetic_score,
|
||||
asethetic_score_siglip_metadata_path=args.asethetic_score_siglip_metadata_path,
|
||||
min_asethetic_score_siglip=args.min_asethetic_score_siglip,
|
||||
text_score_metadata_path=args.text_score_metadata_path,
|
||||
min_text_score=args.min_text_score,
|
||||
motion_score_metadata_path=args.motion_score_metadata_path,
|
||||
min_motion_score=args.min_motion_score,
|
||||
)
|
||||
video_path_list = [os.path.join(args.video_folder, video_path) for video_path in video_path_list]
|
||||
# Sorting to guarantee the same result for each process.
|
||||
video_path_list = natsorted(video_path_list)
|
||||
|
||||
state = PartialState()
|
||||
|
||||
# Accelerate model initialization
|
||||
setattr(torch.nn.Linear, "reset_parameters", lambda self: None)
|
||||
setattr(torch.nn.LayerNorm, "reset_parameters", lambda self: None)
|
||||
torch.nn.init.kaiming_uniform_ = skip
|
||||
torch.nn.init.kaiming_normal_ = skip
|
||||
torch.nn.init.uniform_ = skip
|
||||
torch.nn.init.normal_ = skip
|
||||
|
||||
tokenizer = AutoTokenizer.from_pretrained(os.path.join(args.model_path, "llm"), use_fast=False)
|
||||
tinychat.utils.constants.LLAVA_DEFAULT_IMAGE_PATCH_TOKEN_IDX = (
|
||||
tokenizer.convert_tokens_to_ids(
|
||||
[tinychat.utils.constants.LLAVA_DEFAULT_IMAGE_PATCH_TOKEN]
|
||||
)[0]
|
||||
)
|
||||
config = AutoConfig.from_pretrained(args.model_path, trust_remote_code=True)
|
||||
model = VilaLlamaForCausalLM(config).half()
|
||||
tinychat.utils.constants.LLAVA_DEFAULT_IMAGE_PATCH_TOKEN_IDX = (
|
||||
tokenizer.convert_tokens_to_ids(
|
||||
[tinychat.utils.constants.LLAVA_DEFAULT_IMAGE_PATCH_TOKEN]
|
||||
)[0]
|
||||
)
|
||||
vision_tower = model.get_vision_tower()
|
||||
# if not vision_tower.is_loaded:
|
||||
# vision_tower.load_model()
|
||||
image_processor = vision_tower.image_processor
|
||||
# vision_tower = vision_tower.half()
|
||||
|
||||
if args.precision == "W16A16":
|
||||
pbar = tqdm(range(1))
|
||||
pbar.set_description("Loading checkpoint shards")
|
||||
for i in pbar:
|
||||
model.llm = load_checkpoint_and_dispatch(
|
||||
model.llm,
|
||||
os.path.join(args.model_path, "llm"),
|
||||
no_split_module_classes=[
|
||||
"OPTDecoderLayer",
|
||||
"LlamaDecoderLayer",
|
||||
"BloomBlock",
|
||||
"MPTBlock",
|
||||
"DecoderLayer",
|
||||
"CLIPEncoderLayer",
|
||||
],
|
||||
).to(state.device)
|
||||
model = model.to(state.device)
|
||||
|
||||
elif args.precision == "W4A16":
|
||||
from tinychat.utils.load_quant import load_awq_model
|
||||
# Auto load quant_path from the 3b/8b/13b/40b model.
|
||||
if args.quant_path is None:
|
||||
if "VILA1.5-3b-s2-AWQ" in args.model_path:
|
||||
args.quant_path = os.path.join(args.model_path, "llm/vila-1.5-3b-s2-w4-g128-awq-v2.pt")
|
||||
elif "VILA1.5-3b-AWQ" in args.model_path:
|
||||
args.quant_path = os.path.join(args.model_path, "llm/vila-1.5-3b-w4-g128-awq-v2.pt")
|
||||
elif "Llama-3-VILA1.5-8b-AWQ" in args.model_path:
|
||||
args.quant_path = os.path.join(args.model_path, "llm/llama-3-vila1.5-8b-w4-g128-awq-v2.pt")
|
||||
elif "VILA1.5-13b-AWQ" in args.model_path:
|
||||
args.quant_path = os.path.join(args.model_path, "llm/vila-1.5-13b-w4-g128-awq-v2.pt")
|
||||
elif "VILA1.5-40b-AWQ" in args.model_path:
|
||||
args.quant_path = os.path.join(args.model_path, "llm/vila-1.5-40b-w4-g128-awq-v2.pt")
|
||||
model.llm = load_awq_model(model.llm, args.quant_path, 4, 128, state.device)
|
||||
from tinychat.modules import (
|
||||
make_fused_mlp,
|
||||
make_fused_vision_attn,
|
||||
make_quant_attn,
|
||||
make_quant_norm,
|
||||
)
|
||||
|
||||
make_quant_attn(model.llm, state.device)
|
||||
make_quant_norm(model.llm)
|
||||
# make_fused_mlp(model)
|
||||
# make_fused_vision_attn(model,state.device)
|
||||
model = model.to(state.device)
|
||||
|
||||
else:
|
||||
raise NotImplementedError(f"Precision {args.precision} is not supported.")
|
||||
|
||||
device_warmup(state.device)
|
||||
tune_llava_patch_embedding(vision_tower, device=state.device)
|
||||
|
||||
stream_generator = LlavaStreamGenerator
|
||||
|
||||
model_prompter = get_prompter(
|
||||
args.model_type, args.model_path, False, False
|
||||
)
|
||||
stop_token_ids = get_stop_token_ids(args.model_type, args.model_path)
|
||||
|
||||
model.eval()
|
||||
|
||||
index = len(video_path_list) - len(video_path_list) % state.num_processes
|
||||
# Avoid the NCCL timeout in the final gather operation.
|
||||
logger.info(f"Drop {len(video_path_list) % state.num_processes} videos to ensure each process handles the same number of videos.")
|
||||
video_path_list = video_path_list[:index]
|
||||
logger.info(f"{len(video_path_list)} videos are to be processed.")
|
||||
|
||||
result_dict = {args.video_path_column: [], args.caption_column: []}
|
||||
with state.split_between_processes(video_path_list) as splitted_video_path_list:
|
||||
# TODO: Use VideoDataset.
|
||||
for i, video_path in enumerate(tqdm(splitted_video_path_list)):
|
||||
try:
|
||||
image_list = extract_uniform_frames(video_path, args.num_sampled_frames)
|
||||
image_num = len(image_list)
|
||||
# Similar operation in model_worker.py
|
||||
image_tensor = process_images(image_list, image_processor, model.config)
|
||||
if type(image_tensor) is list:
|
||||
image_tensor = [
|
||||
image.to(state.device, dtype=torch.float16) for image in image_tensor
|
||||
]
|
||||
else:
|
||||
image_tensor = image_tensor.to(state.device, dtype=torch.float16)
|
||||
|
||||
input_prompt = args.input_prompt
|
||||
# Insert image here
|
||||
image_token = get_image_token(model, args.model_path)
|
||||
image_token_holder = tinychat.utils.constants.LLAVA_DEFAULT_IM_TOKEN_PLACE_HOLDER
|
||||
im_token_count = input_prompt.count(image_token_holder)
|
||||
if im_token_count == 0:
|
||||
model_prompter.insert_prompt(image_token * image_num + input_prompt)
|
||||
else:
|
||||
assert im_token_count == image_num
|
||||
input_prompt = input_prompt.replace(image_token_holder, image_token)
|
||||
model_prompter.insert_prompt(input_prompt)
|
||||
output_stream = stream_generator(
|
||||
model,
|
||||
tokenizer,
|
||||
model_prompter.model_input,
|
||||
gen_params,
|
||||
device=state.device,
|
||||
stop_token_ids=stop_token_ids,
|
||||
image_tensor=image_tensor,
|
||||
)
|
||||
outputs = stream_output(output_stream)
|
||||
if len(outputs) != 0:
|
||||
result_dict[args.video_path_column].append(Path(video_path).name)
|
||||
result_dict[args.caption_column].append(outputs)
|
||||
|
||||
except Exception as e:
|
||||
logger.warning(f"VILA with {video_path} failed. Error is {e}.")
|
||||
|
||||
if i != 0 and i % args.saved_freq == 0:
|
||||
state.wait_for_everyone()
|
||||
gathered_result_dict = {k: gather_object(v) for k, v in result_dict.items()}
|
||||
if state.is_main_process and len(gathered_result_dict[args.video_path_column]) != 0:
|
||||
result_df = pd.DataFrame(gathered_result_dict)
|
||||
if args.saved_path.endswith(".csv"):
|
||||
header = False if os.path.exists(args.saved_path) else True
|
||||
result_df.to_csv(args.saved_path, header=header, index=False, mode="a")
|
||||
elif args.saved_path.endswith(".jsonl"):
|
||||
result_df.to_json(args.saved_path, orient="records", lines=True, mode="a", force_ascii=False)
|
||||
logger.info(f"Save result to {args.saved_path}.")
|
||||
for k in result_dict.keys():
|
||||
result_dict[k] = []
|
||||
|
||||
state.wait_for_everyone()
|
||||
gathered_result_dict = {k: gather_object(v) for k, v in result_dict.items()}
|
||||
if state.is_main_process and len(gathered_result_dict[args.video_path_column]) != 0:
|
||||
result_df = pd.DataFrame(gathered_result_dict)
|
||||
if args.saved_path.endswith(".csv"):
|
||||
header = False if os.path.exists(args.saved_path) else True
|
||||
result_df.to_csv(args.saved_path, header=header, index=False, mode="a")
|
||||
elif args.saved_path.endswith(".jsonl"):
|
||||
result_df.to_json(args.saved_path, orient="records", lines=True, mode="a", force_ascii=False)
|
||||
logger.info(f"Save result to {args.saved_path}.")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
args = parse_args()
|
||||
main(args)
|
||||
Reference in New Issue
Block a user