fishxinyu's picture
download
raw
8.02 kB
import torch
import os
from vbench import VBench
from vbench.distributed import dist_init, print0
from datetime import datetime
import argparse
import json
def parse_args():
CUR_DIR = os.path.dirname(os.path.abspath(__file__))
parser = argparse.ArgumentParser(description='VBench', formatter_class=argparse.RawTextHelpFormatter)
parser.add_argument(
"--output_path",
type=str,
default='./evaluation_results/',
help="output path to save the evaluation results",
)
parser.add_argument(
"--full_json_dir",
type=str,
default=f'{CUR_DIR}/vbench/VBench_full_info.json',
help="path to save the json file that contains the prompt and dimension information",
)
parser.add_argument(
"--videos_path",
type=str,
required=True,
help="folder that contains the sampled videos",
)
parser.add_argument(
"--dimension",
nargs='+',
required=True,
help="list of evaluation dimensions, usage: --dimension <dim_1> <dim_2>",
)
parser.add_argument(
"--load_ckpt_from_local",
type=bool,
required=False,
help="whether load checkpoints from local default paths (assuming you have downloaded the checkpoints locally",
)
parser.add_argument(
"--read_frame",
type=bool,
required=False,
help="whether directly read frames, or directly read videos",
)
parser.add_argument(
"--mode",
choices=['custom_input', 'vbench_standard', 'vbench_category'],
default='vbench_standard',
help="""This flags determine the mode of evaluations, choose one of the following:
1. "custom_input": receive input prompt from either --prompt/--prompt_file flags or the filename
2. "vbench_standard": evaluate on standard prompt suite of VBench
3. "vbench_category": evaluate on specific category
""",
)
parser.add_argument(
"--prompt",
type=str,
default="None",
help="""Specify the input prompt
If not specified, filenames will be used as input prompts
* Mutually exclusive to --prompt_file.
** This option must be used with --mode=custom_input flag
"""
)
parser.add_argument(
"--prompt_file",
type=str,
required=False,
help="""Specify the path of the file that contains prompt lists
If not specified, filenames will be used as input prompts
* Mutually exclusive to --prompt.
** This option must be used with --mode=custom_input flag
"""
)
parser.add_argument(
"--category",
type=str,
required=False,
help="""This is for mode=='vbench_category'
The category to evaluate on, usage: --category=animal.
""",
)
## for dimension specific params ###
parser.add_argument(
"--imaging_quality_preprocessing_mode",
type=str,
required=False,
default='longer',
help="""This is for setting preprocessing in imaging_quality
1. 'shorter': if the shorter side is more than 512, the image is resized so that the shorter side is 512.
2. 'longer': if the longer side is more than 512, the image is resized so that the longer side is 512.
3. 'shorter_centercrop': if the shorter side is more than 512, the image is resized so that the shorter side is 512.
Then the center 512 x 512 after resized is used for evaluation.
4. 'None': no preprocessing
""",
)
parser.add_argument(
"--distributed",
action="store_true",
default=False,
help="Initialize torch.distributed (required for multi-GPU runs). Single-GPU runs do not need this.",
)
args = parser.parse_args()
return args
DIMENSION_ABBREV = {
'aesthetic_quality': 'AQ',
'appearance_style': 'AS',
'background_consistency': 'BC',
'color': 'CL',
'dynamic_degree': 'DD',
'human_action': 'HA',
'imaging_quality': 'IQ',
'motion_smoothness': 'MS',
'multiple_objects': 'MO',
'object_class': 'OC',
'overall_consistency': 'OvC',
'scene': 'SC',
'spatial_relationship': 'SR',
'subject_consistency': 'SuC',
'temporal_flickering': 'TF',
'temporal_style': 'TS',
}
def save_results_table(results_path, table_path):
with open(results_path, 'r') as f:
results = json.load(f)
dim_rows = []
for dim, abbrev in sorted(DIMENSION_ABBREV.items(), key=lambda x: x[1]):
if dim not in results:
continue
score = results[dim]
avg = score[0] if isinstance(score, (list, tuple)) else score
per_video = {os.path.basename(e['video_path']): e['video_results']
for e in score[1]} if isinstance(score, (list, tuple)) else {}
dim_rows.append((abbrev, dim, avg, per_video))
if not dim_rows:
return
abbrevs = [r[0] for r in dim_rows]
overall = sum(r[2] for r in dim_rows) / len(dim_rows)
# collect all video names (preserve order from first dim)
all_videos = list(dict.fromkeys(
v for _, _, _, pv in dim_rows for v in pv
))
lines = []
header = 'video,' + ','.join(abbrevs) + ',Avg'
lines.append(header)
avg_scores = [r[2] for r in dim_rows]
lines.append('Avg,' + ','.join(f"{s*100:.2f}" for s in avg_scores) + f",{overall*100:.2f}")
def fmt(v):
if v != v: # nan
return ''
return f"{v*100:.2f}" if 0 <= v <= 1 else f"{v:.2f}"
for vid in all_videos:
scores = [r[3].get(vid, float('nan')) for r in dim_rows]
normalized = [s if 0 <= s <= 1 else float('nan') for s in scores]
valid = [s for s in normalized if s == s]
vid_avg = sum(valid) / len(valid) if valid else float('nan')
lines.append(vid + ',' + ','.join(fmt(s) for s in scores) + f",{vid_avg*100:.2f}")
with open(table_path, 'w') as f:
f.write('\n'.join(lines) + '\n')
print0(f'Results table saved to {table_path}')
def main():
args = parse_args()
if args.distributed:
dist_init()
print0(f'args: {args}')
device = torch.device("cuda")
my_VBench = VBench(device, args.full_json_dir, args.output_path)
print0(f'start evaluation')
current_time = datetime.now().strftime('%Y-%m-%d-%H:%M:%S')
kwargs = {}
prompt = []
if (args.prompt_file is not None) and (args.prompt != "None"):
raise Exception("--prompt_file and --prompt cannot be used together")
if (args.prompt_file is not None or args.prompt != "None") and (args.mode!='custom_input'):
raise Exception("must set --mode=custom_input for using external prompt")
if args.prompt_file:
with open(args.prompt_file, 'r') as f:
prompt = json.load(f)
assert type(prompt) == dict, "Invalid prompt file format. The correct format is {\"video_path\": prompt, ... }"
elif args.prompt != "None":
prompt = [args.prompt]
if args.category != "":
kwargs['category'] = args.category
kwargs['imaging_quality_preprocessing_mode'] = args.imaging_quality_preprocessing_mode
run_name = f'results_{current_time}'
my_VBench.evaluate(
videos_path = args.videos_path,
name = run_name,
prompt_list=prompt, # pass in [] to read prompt from filename
dimension_list = args.dimension,
local=args.load_ckpt_from_local,
read_frame=args.read_frame,
mode=args.mode,
**kwargs
)
results_path = os.path.join(args.output_path, f'{run_name}_eval_results.json')
table_path = os.path.join(args.output_path, f'{run_name}_summary_table.csv')
if os.path.exists(results_path):
save_results_table(results_path, table_path)
print0('done')
if __name__ == "__main__":
main()

Xet Storage Details

Size:
8.02 kB
·
Xet hash:
bef639518ec20ac266e84425a5770928532a92d613e343022d7fd5c52e1d11af

Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.