Add model compiling to the ui
This commit is contained in:
@@ -9,7 +9,7 @@ import {
|
||||
jobTypeOptions,
|
||||
SampleTags,
|
||||
} from './options';
|
||||
import { defaultDatasetConfig } from './jobConfig';
|
||||
import { defaultCompileOptions, defaultDatasetConfig } from './jobConfig';
|
||||
import { GroupedSelectOption, JobConfig, SelectOption } from '@/types';
|
||||
import { objectCopy, tagsToObj, objToTags } from '@/utils/basic';
|
||||
import {
|
||||
@@ -368,7 +368,7 @@ export default function SimpleJob({
|
||||
)}
|
||||
</Card>
|
||||
{disableSections.includes('model.quantize') ? null : (
|
||||
<Card title="Quantization">
|
||||
<Card title="Quantize / Compile">
|
||||
<SelectInput
|
||||
label="Transformer"
|
||||
value={jobConfig.config.process[0].model.quantize ? jobConfig.config.process[0].model.qtype : ''}
|
||||
@@ -401,6 +401,25 @@ export default function SimpleJob({
|
||||
options={quantizationOptions}
|
||||
/>
|
||||
)}
|
||||
<FormGroup label="Compile Options">
|
||||
<></>
|
||||
</FormGroup>
|
||||
<Checkbox
|
||||
label="Compile Model"
|
||||
checked={jobConfig.config.process[0].train.compile || false}
|
||||
onChange={value => {
|
||||
setJobConfig(value, 'config.process[0].train.compile');
|
||||
if (value) {
|
||||
for (const key in defaultCompileOptions) {
|
||||
setJobConfig((defaultCompileOptions as any)[key], `config.process[0].train.${key}`);
|
||||
}
|
||||
} else {
|
||||
for (const key in defaultCompileOptions) {
|
||||
setJobConfig(undefined, `config.process[0].train.${key}`);
|
||||
}
|
||||
}
|
||||
}}
|
||||
/>
|
||||
</Card>
|
||||
)}
|
||||
{modelArch?.additionalSections?.includes('model.multistage') && (
|
||||
|
||||
@@ -31,6 +31,14 @@ export const defaultSliderConfig: SliderConfig = {
|
||||
anchor_class: '',
|
||||
};
|
||||
|
||||
export const defaultCompileOptions = {
|
||||
block_compile: false,
|
||||
compile_mode: 'default',
|
||||
compile_fullgraph: false,
|
||||
compile_dynamic: false,
|
||||
cache_size_limit: undefined,
|
||||
};
|
||||
|
||||
export const defaultJobConfig: JobConfig = {
|
||||
job: 'extension',
|
||||
config: {
|
||||
@@ -94,6 +102,7 @@ export const defaultJobConfig: JobConfig = {
|
||||
diff_output_preservation_class: 'person',
|
||||
switch_boundary_every: 1,
|
||||
loss_type: 'mse',
|
||||
compile: false,
|
||||
},
|
||||
logging: {
|
||||
log_every: 1,
|
||||
|
||||
@@ -151,6 +151,12 @@ export interface TrainConfig {
|
||||
differential_guidance_scale?: number;
|
||||
audio_loss_multiplier?: number;
|
||||
max_loss?: number | null;
|
||||
compile?: boolean;
|
||||
block_compile?: boolean;
|
||||
compile_mode?: 'default' | 'max-autotune' | 'fastest';
|
||||
compile_fullgraph?: boolean;
|
||||
compile_dynamic?: boolean;
|
||||
cache_size_limit?: number;
|
||||
}
|
||||
|
||||
export interface QuantizeKwargsConfig {
|
||||
|
||||
@@ -1 +1 @@
|
||||
VERSION = "0.10.6"
|
||||
VERSION = "0.10.7"
|
||||
|
||||
Reference in New Issue
Block a user