modify ace inference and ace yaml
This commit is contained in:
@@ -156,17 +156,17 @@ class RefinerInference(DiffusionInference):
|
||||
noise.append(noise_)
|
||||
noise, x_shapes = pack_imagelist_into_tensor(noise)
|
||||
if reverse_scale > 0:
|
||||
if self.use_dynamic_model: self.dynamic_load(self.first_stage_model, 'first_stage_model')
|
||||
self.dynamic_load(self.first_stage_model, 'first_stage_model')
|
||||
x_samples = [x.unsqueeze(0) for x in x_samples]
|
||||
x_start = self.encode_first_stage(x_samples, **kwargs)
|
||||
if self.use_dynamic_model: self.dynamic_unload(self.first_stage_model,
|
||||
self.dynamic_unload(self.first_stage_model,
|
||||
'first_stage_model',
|
||||
skip_loaded=True)
|
||||
skip_loaded=not self.use_dynamic_model)
|
||||
x_start, _ = pack_imagelist_into_tensor(x_start)
|
||||
else:
|
||||
x_start = None
|
||||
# cond stage
|
||||
if self.use_dynamic_model: self.dynamic_load(self.cond_stage_model, 'cond_stage_model')
|
||||
self.dynamic_load(self.cond_stage_model, 'cond_stage_model')
|
||||
function_name, dtype = self.get_function_info(self.cond_stage_model)
|
||||
with torch.autocast('cuda',
|
||||
enabled=dtype == 'float16',
|
||||
@@ -174,12 +174,12 @@ class RefinerInference(DiffusionInference):
|
||||
ctx = getattr(get_model(self.cond_stage_model),
|
||||
function_name)(prompt)
|
||||
ctx["x_shapes"] = x_shapes
|
||||
if self.use_dynamic_model: self.dynamic_unload(self.cond_stage_model,
|
||||
self.dynamic_unload(self.cond_stage_model,
|
||||
'cond_stage_model',
|
||||
skip_loaded=True)
|
||||
skip_loaded=not self.use_dynamic_model)
|
||||
|
||||
|
||||
if self.use_dynamic_model: self.dynamic_load(self.diffusion_model, 'diffusion_model')
|
||||
self.dynamic_load(self.diffusion_model, 'diffusion_model')
|
||||
# UNet use input n_prompt
|
||||
function_name, dtype = self.get_function_info(
|
||||
self.diffusion_model)
|
||||
@@ -207,14 +207,14 @@ class RefinerInference(DiffusionInference):
|
||||
x=x_start,
|
||||
**kwargs).float()
|
||||
latent = unpack_tensor_into_imagelist(latent, x_shapes)
|
||||
if self.use_dynamic_model: self.dynamic_unload(self.diffusion_model,
|
||||
self.dynamic_unload(self.diffusion_model,
|
||||
'diffusion_model',
|
||||
skip_loaded=True)
|
||||
if self.use_dynamic_model: self.dynamic_load(self.first_stage_model, 'first_stage_model')
|
||||
skip_loaded=not self.use_dynamic_model)
|
||||
self.dynamic_load(self.first_stage_model, 'first_stage_model')
|
||||
x_samples = self.decode_first_stage(latent)
|
||||
if self.use_dynamic_model: self.dynamic_unload(self.first_stage_model,
|
||||
self.dynamic_unload(self.first_stage_model,
|
||||
'first_stage_model',
|
||||
skip_loaded=True)
|
||||
skip_loaded=not self.use_dynamic_model)
|
||||
return x_samples
|
||||
|
||||
|
||||
@@ -398,11 +398,11 @@ class ACEInference(DiffusionInference):
|
||||
if use_ace and (not is_txt_image or refiner_scale <= 0):
|
||||
ctx, null_ctx = {}, {}
|
||||
# Get Noise Shape
|
||||
if self.use_dynamic_model: self.dynamic_load(self.first_stage_model, 'first_stage_model')
|
||||
self.dynamic_load(self.first_stage_model, 'first_stage_model')
|
||||
x = self.encode_first_stage(image)
|
||||
if self.use_dynamic_model: self.dynamic_unload(self.first_stage_model,
|
||||
self.dynamic_unload(self.first_stage_model,
|
||||
'first_stage_model',
|
||||
skip_loaded=True)
|
||||
skip_loaded=not self.use_dynamic_model)
|
||||
noise = [
|
||||
torch.empty(*i.shape, device=we.device_id).normal_(generator=g)
|
||||
for i in x
|
||||
@@ -416,7 +416,7 @@ class ACEInference(DiffusionInference):
|
||||
ctx['x_mask'] = null_ctx['x_mask'] = cond_mask
|
||||
|
||||
# Encode Prompt
|
||||
if self.use_dynamic_model: self.dynamic_load(self.cond_stage_model, 'cond_stage_model')
|
||||
self.dynamic_load(self.cond_stage_model, 'cond_stage_model')
|
||||
function_name, dtype = self.get_function_info(self.cond_stage_model)
|
||||
cont, cont_mask = getattr(get_model(self.cond_stage_model),
|
||||
function_name)(prompt)
|
||||
@@ -426,14 +426,14 @@ class ACEInference(DiffusionInference):
|
||||
function_name)(n_prompt)
|
||||
null_cont, null_cont_mask = self.cond_stage_embeddings(
|
||||
prompt, edit_image, null_cont, null_cont_mask)
|
||||
if self.use_dynamic_model: self.dynamic_unload(self.cond_stage_model,
|
||||
self.dynamic_unload(self.cond_stage_model,
|
||||
'cond_stage_model',
|
||||
skip_loaded=False)
|
||||
skip_loaded=not self.use_dynamic_model)
|
||||
ctx['crossattn'] = cont
|
||||
null_ctx['crossattn'] = null_cont
|
||||
|
||||
# Encode Edit Images
|
||||
if self.use_dynamic_model: self.dynamic_load(self.first_stage_model, 'first_stage_model')
|
||||
self.dynamic_load(self.first_stage_model, 'first_stage_model')
|
||||
edit_image = [to_device(i, strict=False) for i in edit_image]
|
||||
edit_image_mask = [to_device(i, strict=False) for i in edit_image_mask]
|
||||
e_img, e_mask = [], []
|
||||
@@ -444,14 +444,14 @@ class ACEInference(DiffusionInference):
|
||||
m = [None] * len(u)
|
||||
e_img.append(self.encode_first_stage(u, **kwargs))
|
||||
e_mask.append([self.interpolate_func(i) for i in m])
|
||||
if self.use_dynamic_model: self.dynamic_unload(self.first_stage_model,
|
||||
self.dynamic_unload(self.first_stage_model,
|
||||
'first_stage_model',
|
||||
skip_loaded=True)
|
||||
skip_loaded=not self.use_dynamic_model)
|
||||
null_ctx['edit'] = ctx['edit'] = e_img
|
||||
null_ctx['edit_mask'] = ctx['edit_mask'] = e_mask
|
||||
|
||||
# Diffusion Process
|
||||
if self.use_dynamic_model: self.dynamic_load(self.diffusion_model, 'diffusion_model')
|
||||
self.dynamic_load(self.diffusion_model, 'diffusion_model')
|
||||
function_name, dtype = self.get_function_info(self.diffusion_model)
|
||||
with torch.autocast('cuda',
|
||||
enabled=dtype in ('float16', 'bfloat16'),
|
||||
@@ -492,17 +492,17 @@ class ACEInference(DiffusionInference):
|
||||
guide_rescale=guide_rescale,
|
||||
return_intermediate=None,
|
||||
**kwargs)
|
||||
if self.use_dynamic_model: self.dynamic_unload(self.diffusion_model,
|
||||
self.dynamic_unload(self.diffusion_model,
|
||||
'diffusion_model',
|
||||
skip_loaded=False)
|
||||
skip_loaded=not self.use_dynamic_model)
|
||||
|
||||
# Decode to Pixel Space
|
||||
if self.use_dynamic_model: self.dynamic_load(self.first_stage_model, 'first_stage_model')
|
||||
self.dynamic_load(self.first_stage_model, 'first_stage_model')
|
||||
samples = unpack_tensor_into_imagelist(latent, x_shapes)
|
||||
x_samples = self.decode_first_stage(samples)
|
||||
if self.use_dynamic_model: self.dynamic_unload(self.first_stage_model,
|
||||
self.dynamic_unload(self.first_stage_model,
|
||||
'first_stage_model',
|
||||
skip_loaded=False)
|
||||
skip_loaded=not self.use_dynamic_model)
|
||||
x_samples = [x.squeeze(0) for x in x_samples]
|
||||
else:
|
||||
x_samples = image
|
||||
|
||||
@@ -197,7 +197,7 @@ class LatentDiffusionSolver(BaseSolver):
|
||||
else:
|
||||
self.logger.info('Use default backend.')
|
||||
self.use_scaler = cfg.get('USE_SCALER', True)
|
||||
self.enable_gradscaler = cfg.get('ENABLE_GRADSCALER', True)
|
||||
self.enable_gradscaler = cfg.get('ENABLE_GRADSCALER', False)
|
||||
self.use_orig_params = cfg.get('USE_ORIG_PARAMS', False)
|
||||
self.model_shard = cfg.get('SHARDING_STRATEGY', 'full_shard')
|
||||
self.sharding_size = cfg.get('SHARDING_SIZE', None)
|
||||
@@ -422,8 +422,10 @@ class LatentDiffusionSolver(BaseSolver):
|
||||
process_group=None)
|
||||
else:
|
||||
self.scaler = amp.GradScaler(enabled=self.enable_gradscaler)
|
||||
else:
|
||||
elif self.cfg.DTYPE in ['float16']:
|
||||
self.scaler = amp.GradScaler()
|
||||
else:
|
||||
self.scaler = None
|
||||
else:
|
||||
self.scaler = None
|
||||
|
||||
|
||||
@@ -1,6 +1,10 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
import copy
|
||||
from enum import Enum
|
||||
import os
|
||||
|
||||
from scepter.modules.utils.file_system import FS
|
||||
|
||||
|
||||
class Media(Enum):
|
||||
@@ -13,15 +17,17 @@ class Media(Enum):
|
||||
|
||||
|
||||
class HtmlVisualization(object):
|
||||
def __init__(self,
|
||||
allow_annotation=False,
|
||||
slice_size=1000,
|
||||
align='center',
|
||||
width_scale='60%',
|
||||
title='Visualization',
|
||||
height=600,
|
||||
width=None,
|
||||
text_cols=40):
|
||||
def __init__(
|
||||
self,
|
||||
allow_annotation=False,
|
||||
slice_size=1000,
|
||||
align='center',
|
||||
width_scale='60%',
|
||||
title='Visualization',
|
||||
height=600,
|
||||
width=None,
|
||||
text_cols=40
|
||||
):
|
||||
self.content_list = []
|
||||
self.rows_meta = []
|
||||
self.allow_annotation = allow_annotation
|
||||
@@ -31,9 +37,9 @@ class HtmlVisualization(object):
|
||||
self.title = title
|
||||
self.html_start = '<html>'
|
||||
self.html_head = f'<head><meta charset="utf-8"><title>{title}</title></head>'
|
||||
self.height = height if height is not None else '600'
|
||||
self.width = width if width is not None else 'auto'
|
||||
self.text_cols = text_cols if text_cols is not None else 'auto'
|
||||
self.height = height if height is not None else "600"
|
||||
self.width = width if width is not None else "auto"
|
||||
self.text_cols = text_cols if text_cols is not None else "auto"
|
||||
self.html_style = ('''
|
||||
<style> \n
|
||||
.container {
|
||||
@@ -82,10 +88,12 @@ class HtmlVisualization(object):
|
||||
resize: none; \n
|
||||
border: 1px solid #ccc; \n
|
||||
} \n
|
||||
.large-checkbox {transform: scale(2.5); margin-left: 20px; margin-bottom: 20px; vertical-align: middle;} \n
|
||||
</style> \n
|
||||
\n
|
||||
'''.replace('{width_scale}', self.width_scale).replace(
|
||||
'{align}', self.align).replace('{pair_height}', f'{self.height}'))
|
||||
'''.replace('{width_scale}',
|
||||
self.width_scale).replace('{align}', self.align)
|
||||
.replace('{pair_height}', f'{self.height}'))
|
||||
|
||||
self.html_body_script = '''
|
||||
<script>\n
|
||||
@@ -99,6 +107,7 @@ class HtmlVisualization(object):
|
||||
isDragging = true;\n
|
||||
});\n
|
||||
|
||||
|
||||
container.addEventListener('mouseup', () => {\n
|
||||
isDragging = true;\n
|
||||
});\n
|
||||
@@ -111,6 +120,9 @@ class HtmlVisualization(object):
|
||||
|
||||
let percentage = (clientX - left) / width * 100;\n
|
||||
|
||||
|
||||
// 限制百分比在0到100之间\n
|
||||
|
||||
percentage = Math.max(0, Math.min(100, percentage));\n
|
||||
|
||||
media2.style.clipPath = `inset(0 ${100 - percentage}% 0 0)`;\n
|
||||
@@ -120,6 +132,7 @@ class HtmlVisualization(object):
|
||||
console.info(slider.style.left);\n
|
||||
|
||||
});\n
|
||||
// 初始化滑块位置\n
|
||||
slider.style.left = '50%';\n
|
||||
});\n
|
||||
</script>\n
|
||||
@@ -152,16 +165,17 @@ class HtmlVisualization(object):
|
||||
|
||||
'''
|
||||
self.label_button = (
|
||||
'<table><tr><td>' +
|
||||
"<button style='height: 50px;' type=\"button\" onclick=\"saveSamples()\">Save Samples</button>"
|
||||
+ '</td></tr></table>')
|
||||
'<table><tr><td>' +
|
||||
"<button style='height: 50px;' type=\"button\" onclick=\"saveSamples()\">Save Samples</button>"
|
||||
+ '</td></tr></table>')
|
||||
|
||||
def format_col(self,
|
||||
content='',
|
||||
label='',
|
||||
type=Media.TEXT,
|
||||
show_label=True,
|
||||
cols_span=1):
|
||||
cols_span=1
|
||||
):
|
||||
if type == Media.TEXT:
|
||||
ret_str = '<textarea' # noqa: E501
|
||||
if self.height is not None:
|
||||
@@ -171,7 +185,7 @@ class HtmlVisualization(object):
|
||||
cols = f"cols={self.text_cols * cols_span}"
|
||||
ret_str += f" {cols}"
|
||||
ret_str += f'>"{content}"</textarea>'
|
||||
sec_ret_str = f'<font size="3"><strong>{label}<strong></font>' if show_label else ''
|
||||
sec_ret_str = f'<font size="3"><strong>{label}<strong></font>' if show_label else ""
|
||||
elif type == Media.IMAGE:
|
||||
ret_str = f'<img src="{content}"'
|
||||
if self.height is not None:
|
||||
@@ -181,7 +195,7 @@ class HtmlVisualization(object):
|
||||
width = f'width="{self.width}"'
|
||||
ret_str += f" {width}"
|
||||
ret_str += ' >'
|
||||
sec_ret_str = f'<font size="3"><strong>{label}<strong></font>' if show_label else ''
|
||||
sec_ret_str = f'<font size="3"><strong>{label}<strong></font>' if show_label else ""
|
||||
elif type == Media.VIDEO:
|
||||
ret_str = '<video' # noqa
|
||||
if self.height is not None:
|
||||
@@ -192,70 +206,85 @@ class HtmlVisualization(object):
|
||||
ret_str += f" {width}"
|
||||
ret_str += ' preload="none" autoplay muted loop>'
|
||||
ret_str += f'<source src="{content}" type="video/mp4"></video>'
|
||||
sec_ret_str = f'<font size="3"><strong>{label}<strong></font>' if show_label else ''
|
||||
sec_ret_str = f'<font size="3"><strong>{label}<strong></font>' if show_label else ""
|
||||
elif type == Media.AUDIO:
|
||||
ret_str = f'<audio src="{content}" controls>'
|
||||
sec_ret_str = f'<font size="3"><strong>{label}<strong></font>' if show_label else ''
|
||||
sec_ret_str = f'<font size="3"><strong>{label}<strong></font>' if show_label else ""
|
||||
elif type == Media.IMAGE_PAIR:
|
||||
assert isinstance(content, (list, tuple)) and len(content) == 2
|
||||
ret_str = f'\n' # noqa
|
||||
ret_str += f' <div class="container"' # noqa
|
||||
ret_str += (
|
||||
f'> \n'
|
||||
f' <div class="image" id="media1">' # noqa
|
||||
f' <img src="{content[1]}" alt="before">\n' # noqa
|
||||
f' </div>\n' # noqa
|
||||
f' <div class="image" id="media2" style="clip-path: inset(0 50% 0 0);">\n' # noqa
|
||||
f' <img src="{content[0]}" alt="after">\n' # noqa
|
||||
f' </div>\n' # noqa
|
||||
f' <div class="slider" id="slider"></div>\n' # noqa
|
||||
f'')
|
||||
sec_ret_str = f'<font size="3"><strong>{label}<strong></font>' if show_label else ''
|
||||
ret_str = f'\n'
|
||||
ret_str += f' <div class="container"'
|
||||
ret_str += (f'> \n'
|
||||
f' <div class="image" id="media1">'
|
||||
f' <img src="{content[1]}" alt="before">\n'
|
||||
f' </div>\n'
|
||||
f' <div class="image" id="media2" style="clip-path: inset(0 50% 0 0);">\n'
|
||||
f' <img src="{content[0]}" alt="after">\n'
|
||||
f' </div>\n'
|
||||
f' <div class="slider" id="slider"></div>\n'
|
||||
f'')
|
||||
sec_ret_str = f'<font size="3"><strong>{label}<strong></font>' if show_label else ""
|
||||
elif type == Media.VIDEO_PAIR:
|
||||
assert isinstance(content, (list, tuple)) and len(content) == 2
|
||||
ret_str = f'\n' # noqa
|
||||
ret_str += f' <div class="container"' # noqa
|
||||
ret_str += (
|
||||
f'> \n'
|
||||
f' <video autoplay muted loop class="video" id="media1"><source src="{content[1]}" type="video/mp4"></video>\n' # noqa
|
||||
f' <video autoplay muted loop class="video" id="media2" style="clip-path: inset(0 50% 0 0);"><source src="{content[0]}" type="video/mp4"></video>\n' # noqa
|
||||
f' <div class="slider" id="slider"></div>\n' # noqa
|
||||
f'</div>')
|
||||
sec_ret_str = f'<font size="3"><strong>{label}<strong></font>' if show_label else ''
|
||||
ret_str = f'\n'
|
||||
ret_str += f' <div class="container"'
|
||||
ret_str += (f'> \n'
|
||||
f' <video autoplay muted loop class="video" id="media1"><source src="{content[1]}" type="video/mp4"></video>\n'
|
||||
f' <video autoplay muted loop class="video" id="media2" style="clip-path: inset(0 50% 0 0);"><source src="{content[0]}" type="video/mp4"></video>\n'
|
||||
f' <div class="slider" id="slider"></div>\n'
|
||||
f'</div>')
|
||||
sec_ret_str = f'<font size="3"><strong>{label}<strong></font>' if show_label else ""
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
if self.allow_annotation:
|
||||
ret_str = f'<label for="sample#sample_id#">{ret_str}</label>'
|
||||
|
||||
if cols_span > 1:
|
||||
ret_str = f'<th colspan="{cols_span}">{ret_str}</th>\n'
|
||||
sec_ret_str = f'<th colspan="{cols_span}">{sec_ret_str}</th>\n' if not sec_ret_str == '' else sec_ret_str
|
||||
sec_ret_str = f'<th colspan="{cols_span}">{sec_ret_str}</th>\n' if not sec_ret_str == "" else sec_ret_str
|
||||
else:
|
||||
ret_str = f'<td>{ret_str}</td>\n'
|
||||
sec_ret_str = f'<td align="center">{sec_ret_str}</td>\n' if not sec_ret_str == '' else sec_ret_str
|
||||
sec_ret_str = f'<td align="center">{sec_ret_str}</td>\n' if not sec_ret_str == "" else sec_ret_str
|
||||
return [ret_str, sec_ret_str]
|
||||
|
||||
def format_row(self):
|
||||
sample_id = 0
|
||||
all_sample_html = []
|
||||
for one_content, one_row_meta in zip(self.content_list,
|
||||
self.rows_meta):
|
||||
one_row_str = '<tr>'
|
||||
one_row_str += '\n'.join([v[0] for v in one_content])
|
||||
if self.allow_annotation:
|
||||
row_meta = '#;#'.join(one_row_meta)
|
||||
one_row_str += (
|
||||
f'<td><input type="checkbox" class="large-checkbox" '
|
||||
f'id="sample{sample_id}" name="sample[]" value="{row_meta}"></td>\n'
|
||||
)
|
||||
one_row_str += '</tr><tr>'
|
||||
one_row_str += '\n'.join([v[1] for v in one_content]) # noqa
|
||||
if self.allow_annotation: # noqa
|
||||
one_row_str += f'<td></td>\n' # noqa
|
||||
one_row_str += '</tr>'
|
||||
if self.allow_annotation:
|
||||
one_row_str = f'<label for="sample{sample_id}">{one_row_str}</label>'
|
||||
all_sample_html.append(one_row_str)
|
||||
sample_id += 1
|
||||
|
||||
return '<table>' + '\n'.join(all_sample_html) + '</table>'
|
||||
all_sample_html = []
|
||||
current_content_list = copy.deepcopy(self.content_list)
|
||||
current_rows_meta = copy.deepcopy(self.rows_meta)
|
||||
|
||||
while len(current_content_list) > 0:
|
||||
sample_id = 0
|
||||
batch_content_list = current_content_list[:self.slice_size]
|
||||
current_content_list = current_content_list[self.slice_size:]
|
||||
batch_rows_meta = current_rows_meta[:self.slice_size]
|
||||
current_rows_meta = current_rows_meta[self.slice_size:]
|
||||
current_sample_html = []
|
||||
for one_content, one_row_meta in zip(batch_content_list,
|
||||
batch_rows_meta):
|
||||
one_row_str = '<tr>'
|
||||
if not self.allow_annotation:
|
||||
one_row_str += '\n'.join([v[0] for v in one_content])
|
||||
else:
|
||||
one_row_str += '\n'.join([v[0].replace('#sample_id#', f'{sample_id}') for v in one_content])
|
||||
row_meta = '#;#'.join(one_row_meta)
|
||||
one_row_str += (
|
||||
f'<td><input type="checkbox" class="large-checkbox" '
|
||||
f'id="sample{sample_id}" name="sample[]" value="{row_meta}"></td>\n'
|
||||
)
|
||||
one_row_str += '</tr><tr>'
|
||||
one_row_str += '\n'.join([v[1] for v in one_content]) # noqa
|
||||
if self.allow_annotation: # noqa
|
||||
one_row_str += f'<td></td>\n' # noqa
|
||||
one_row_str += '</tr>'
|
||||
# if self.allow_annotation:
|
||||
# one_row_str = f'<label for="sample{sample_id}">{one_row_str}</label>'
|
||||
current_sample_html.append(one_row_str)
|
||||
sample_id += 1
|
||||
all_sample_html.append("<table>" + '\n'.join(current_sample_html) + "</table>")
|
||||
|
||||
return all_sample_html
|
||||
|
||||
def add_record(self,
|
||||
content,
|
||||
@@ -275,11 +304,8 @@ class HtmlVisualization(object):
|
||||
if col_id > len(self.content_list[row_id]):
|
||||
raise RuntimeError(
|
||||
'col_id should be next number of the last col_id.')
|
||||
format_col = self.format_col(content,
|
||||
f"{row_id}-{col_id}: {label}",
|
||||
type,
|
||||
show_label=show_label,
|
||||
cols_span=cols_span)
|
||||
format_col = self.format_col(content, f"{row_id}-{col_id}: {label}",
|
||||
type, show_label=show_label, cols_span=cols_span)
|
||||
|
||||
annotation_meta = annotation_meta if annotation_meta else ''
|
||||
if col_id == len(self.content_list[row_id]):
|
||||
@@ -291,14 +317,30 @@ class HtmlVisualization(object):
|
||||
|
||||
def save_html(self, path):
|
||||
html_body = self.format_row()
|
||||
ret_html_list = [
|
||||
self.html_start, self.html_head, self.html_style,
|
||||
self.html_body.replace('{BODY}', html_body)
|
||||
]
|
||||
if self.allow_annotation:
|
||||
ret_html_list.append(self.label_button)
|
||||
ret_html_list.append(self.html_script)
|
||||
ret_html_list.append(self.html_end)
|
||||
ret_html = '\n'.join(ret_html_list)
|
||||
with open(path, 'w') as f:
|
||||
f.write(ret_html)
|
||||
if isinstance(html_body, list) and len(html_body) > 1:
|
||||
try:
|
||||
os.makedirs(path, exist_ok=True)
|
||||
except:
|
||||
print("Create folder path failed.")
|
||||
for html_id, one_html in enumerate(html_body):
|
||||
ret_html_list = [
|
||||
self.html_start, self.html_head, self.html_style,
|
||||
self.html_body.replace('{BODY}', one_html)
|
||||
]
|
||||
if self.allow_annotation:
|
||||
ret_html_list.append(self.label_button)
|
||||
ret_html_list.append(self.html_script)
|
||||
ret_html_list.append(self.html_end)
|
||||
ret_html = '\n'.join(ret_html_list)
|
||||
FS.put_object(ret_html.encode(), os.path.join(path, f"{html_id}.html"))
|
||||
else:
|
||||
ret_html_list = [
|
||||
self.html_start, self.html_head, self.html_style,
|
||||
self.html_body.replace('{BODY}', html_body[0])
|
||||
]
|
||||
if self.allow_annotation:
|
||||
ret_html_list.append(self.label_button)
|
||||
ret_html_list.append(self.html_script)
|
||||
ret_html_list.append(self.html_end)
|
||||
ret_html = '\n'.join(ret_html_list)
|
||||
FS.put_object(ret_html.encode(), path)
|
||||
|
||||
Reference in New Issue
Block a user