Add multi-prompt sequence generation to demo
This commit is contained in:
parent
e55a42aaa3
commit
ff811e49de
220
demo/main.go
220
demo/main.go
@ -5,11 +5,13 @@ package main
|
||||
import (
|
||||
"crypto/rand"
|
||||
"embed"
|
||||
"encoding/binary"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"flag"
|
||||
"fmt"
|
||||
"log"
|
||||
"math"
|
||||
"net/http"
|
||||
"os"
|
||||
"os/exec"
|
||||
@ -30,16 +32,23 @@ var modelUI []byte
|
||||
var localAILogo []byte
|
||||
|
||||
type animation struct {
|
||||
ID string `json:"id"`
|
||||
Prompt string `json:"prompt"`
|
||||
Frames int `json:"frames"`
|
||||
DiffusionSteps int `json:"diffusion_steps"`
|
||||
Seed uint64 `json:"seed"`
|
||||
CreatedAt string `json:"created_at"`
|
||||
Status string `json:"status"`
|
||||
Error string `json:"error,omitempty"`
|
||||
Kind string `json:"kind"`
|
||||
Model string `json:"model"`
|
||||
ID string `json:"id"`
|
||||
Prompt string `json:"prompt"`
|
||||
Frames int `json:"frames"`
|
||||
DiffusionSteps int `json:"diffusion_steps"`
|
||||
Seed uint64 `json:"seed"`
|
||||
CreatedAt string `json:"created_at"`
|
||||
Status string `json:"status"`
|
||||
Error string `json:"error,omitempty"`
|
||||
Kind string `json:"kind"`
|
||||
Model string `json:"model"`
|
||||
Segments []promptSegment `json:"segments,omitempty"`
|
||||
TransitionFrames int `json:"transition_frames,omitempty"`
|
||||
Progress string `json:"progress,omitempty"`
|
||||
}
|
||||
type promptSegment struct {
|
||||
Prompt string `json:"prompt"`
|
||||
Frames int `json:"frames"`
|
||||
}
|
||||
type motionModel struct {
|
||||
ID string `json:"id"`
|
||||
@ -84,6 +93,111 @@ func (g *gallery) list() []*animation {
|
||||
sort.Slice(result, func(i, j int) bool { return result[i].CreatedAt > result[j].CreatedAt })
|
||||
return result
|
||||
}
|
||||
|
||||
func copyFile(dst, src string) error {
|
||||
b, err := os.ReadFile(src)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return os.WriteFile(dst, b, 0600)
|
||||
}
|
||||
|
||||
func readF32(path string) ([]float32, error) {
|
||||
b, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if len(b)%4 != 0 {
|
||||
return nil, fmt.Errorf("invalid F32 file: %s", path)
|
||||
}
|
||||
values := make([]float32, len(b)/4)
|
||||
for i := range values {
|
||||
values[i] = math.Float32frombits(binary.LittleEndian.Uint32(b[i*4:]))
|
||||
}
|
||||
return values, nil
|
||||
}
|
||||
|
||||
func writeF32(path string, values []float32) error {
|
||||
b := make([]byte, len(values)*4)
|
||||
for i, value := range values {
|
||||
binary.LittleEndian.PutUint32(b[i*4:], math.Float32bits(value))
|
||||
}
|
||||
return os.WriteFile(path, b, 0600)
|
||||
}
|
||||
|
||||
func blendQuaternion(a, b []float32, alpha float32) {
|
||||
dot := a[0]*b[0] + a[1]*b[1] + a[2]*b[2] + a[3]*b[3]
|
||||
if dot < 0 {
|
||||
for i := range b {
|
||||
b[i] = -b[i]
|
||||
}
|
||||
}
|
||||
length := float32(0)
|
||||
for i := range a {
|
||||
a[i] = alpha*a[i] + (1-alpha)*b[i]
|
||||
length += a[i] * a[i]
|
||||
}
|
||||
if length > 0 {
|
||||
length = 1 / float32(math.Sqrt(float64(length)))
|
||||
for i := range a {
|
||||
a[i] *= length
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// stitchSegments joins independently sampled demo segments. The overlap is
|
||||
// blended in root space and by normalized linear interpolation for quaternions.
|
||||
// Native observed-motion conditioning is deliberately a later parity step.
|
||||
func stitchSegments(output string, dirs []string, overlap int) error {
|
||||
var roots, rotations []float32
|
||||
for index, dir := range dirs {
|
||||
root, err := readF32(filepath.Join(dir, "root_positions.f32"))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
rot, err := readF32(filepath.Join(dir, "local_rotations_xyzw.f32"))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
frames := len(root) / 3
|
||||
if frames == 0 || len(rot) != frames*22*4 {
|
||||
return fmt.Errorf("invalid motion segment %d", index+1)
|
||||
}
|
||||
if index == 0 {
|
||||
roots, rotations = root, rot
|
||||
continue
|
||||
}
|
||||
n := overlap
|
||||
if n > frames {
|
||||
n = frames
|
||||
}
|
||||
if n > len(roots)/3 {
|
||||
n = len(roots) / 3
|
||||
}
|
||||
for frame := 0; frame < n; frame++ {
|
||||
alpha := float32(0.5)
|
||||
if n > 1 {
|
||||
alpha = 1 - float32(frame)/float32(n-1)
|
||||
}
|
||||
old := (len(roots)/3 - n + frame) * 3
|
||||
newest := frame * 3
|
||||
for axis := 0; axis < 3; axis++ {
|
||||
roots[old+axis] = alpha*roots[old+axis] + (1-alpha)*root[newest+axis]
|
||||
}
|
||||
for joint := 0; joint < 22; joint++ {
|
||||
oldQ := (len(rotations)/4 - n*22 + frame*22 + joint) * 4
|
||||
newQ := (frame*22 + joint) * 4
|
||||
blendQuaternion(rotations[oldQ:oldQ+4], append([]float32(nil), rot[newQ:newQ+4]...), alpha)
|
||||
}
|
||||
}
|
||||
roots = append(roots, root[n*3:]...)
|
||||
rotations = append(rotations, rot[n*22*4:]...)
|
||||
}
|
||||
if err := writeF32(filepath.Join(output, "root_positions.f32"), roots); err != nil {
|
||||
return err
|
||||
}
|
||||
return writeF32(filepath.Join(output, "local_rotations_xyzw.f32"), rotations)
|
||||
}
|
||||
func (g *gallery) worker() {
|
||||
for id := range g.queue {
|
||||
g.mu.Lock()
|
||||
@ -101,11 +215,43 @@ func (g *gallery) worker() {
|
||||
if !ok || !model.Available {
|
||||
err = fmt.Errorf("model %q is not available", item.Model)
|
||||
} else {
|
||||
cmd := exec.Command(g.generator, model.Motion, g.text, filepath.Join(dir, "prompt.txt"), fmt.Sprint(item.Frames), fmt.Sprint(item.DiffusionSteps), fmt.Sprint(item.Seed), dir)
|
||||
cmd.Env = append(os.Environ(), "KIMODO_BACKEND=vulkan")
|
||||
output, runErr := cmd.CombinedOutput()
|
||||
if runErr != nil {
|
||||
err = fmt.Errorf("%w: %s", runErr, strings.TrimSpace(string(output)))
|
||||
segments := item.Segments
|
||||
if len(segments) == 0 {
|
||||
segments = []promptSegment{{Prompt: item.Prompt, Frames: item.Frames}}
|
||||
}
|
||||
segmentDirs := make([]string, 0, len(segments))
|
||||
for index, segment := range segments {
|
||||
g.mu.Lock()
|
||||
item.Progress = fmt.Sprintf("Generating segment %d of %d", index+1, len(segments))
|
||||
_ = g.save(item)
|
||||
g.mu.Unlock()
|
||||
segmentDir := filepath.Join(dir, fmt.Sprintf("segment-%02d", index+1))
|
||||
if err = os.MkdirAll(segmentDir, 0755); err != nil {
|
||||
break
|
||||
}
|
||||
promptPath := filepath.Join(segmentDir, "prompt.txt")
|
||||
if err = os.WriteFile(promptPath, []byte(segment.Prompt), 0600); err != nil {
|
||||
break
|
||||
}
|
||||
cmd := exec.Command(g.generator, model.Motion, g.text, promptPath, fmt.Sprint(segment.Frames), fmt.Sprint(item.DiffusionSteps), fmt.Sprint(item.Seed+uint64(index)), segmentDir)
|
||||
cmd.Env = append(os.Environ(), "KIMODO_BACKEND=vulkan")
|
||||
output, runErr := cmd.CombinedOutput()
|
||||
if runErr != nil {
|
||||
err = fmt.Errorf("segment %d: %w: %s", index+1, runErr, strings.TrimSpace(string(output)))
|
||||
break
|
||||
}
|
||||
segmentDirs = append(segmentDirs, segmentDir)
|
||||
}
|
||||
if err == nil && len(segmentDirs) > 1 {
|
||||
err = stitchSegments(dir, segmentDirs, item.TransitionFrames)
|
||||
}
|
||||
if err == nil && len(segmentDirs) == 1 {
|
||||
for _, name := range []string{"root_positions.f32", "local_rotations_xyzw.f32"} {
|
||||
err = copyFile(filepath.Join(dir, name), filepath.Join(segmentDirs[0], name))
|
||||
if err != nil {
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@ -115,6 +261,7 @@ func (g *gallery) worker() {
|
||||
item.Error = err.Error()
|
||||
} else {
|
||||
item.Status = "ready"
|
||||
item.Progress = ""
|
||||
}
|
||||
if saveErr := g.save(item); saveErr != nil {
|
||||
log.Printf("save %s: %v", item.ID, saveErr)
|
||||
@ -196,19 +343,24 @@ func main() {
|
||||
return
|
||||
}
|
||||
var request struct {
|
||||
Prompt string `json:"prompt"`
|
||||
Frames int `json:"frames"`
|
||||
Steps int `json:"steps"`
|
||||
Seed uint64 `json:"seed"`
|
||||
Model string `json:"model"`
|
||||
Prompt string `json:"prompt"`
|
||||
Segments []promptSegment `json:"segments"`
|
||||
TransitionFrames int `json:"transition_frames"`
|
||||
Frames int `json:"frames"`
|
||||
Steps int `json:"steps"`
|
||||
Seed uint64 `json:"seed"`
|
||||
Model string `json:"model"`
|
||||
}
|
||||
if err := json.NewDecoder(http.MaxBytesReader(w, r.Body, 32<<10)).Decode(&request); err != nil {
|
||||
http.Error(w, "invalid JSON", 400)
|
||||
return
|
||||
}
|
||||
request.Prompt = strings.TrimSpace(request.Prompt)
|
||||
if request.Prompt == "" || len(request.Prompt) > 4096 {
|
||||
http.Error(w, "prompt must be 1..4096 bytes", 400)
|
||||
if len(request.Segments) == 0 {
|
||||
request.Segments = []promptSegment{{Prompt: request.Prompt, Frames: request.Frames}}
|
||||
}
|
||||
if len(request.Segments) > 16 {
|
||||
http.Error(w, "at most 16 prompt segments", 400)
|
||||
return
|
||||
}
|
||||
if request.Frames == 0 {
|
||||
@ -217,8 +369,21 @@ func main() {
|
||||
if request.Steps == 0 {
|
||||
request.Steps = 100
|
||||
}
|
||||
if request.Frames < 1 || request.Frames > 1000 || request.Steps < 1 || request.Steps > 1000 {
|
||||
http.Error(w, "frames and steps must be 1..1000", 400)
|
||||
for index := range request.Segments {
|
||||
request.Segments[index].Prompt = strings.TrimSpace(request.Segments[index].Prompt)
|
||||
if request.Segments[index].Frames == 0 {
|
||||
request.Segments[index].Frames = 150
|
||||
}
|
||||
if request.Segments[index].Prompt == "" || len(request.Segments[index].Prompt) > 4096 || request.Segments[index].Frames < 60 || request.Segments[index].Frames > 300 {
|
||||
http.Error(w, "each prompt segment must be 60..300 frames and 1..4096 bytes", 400)
|
||||
return
|
||||
}
|
||||
}
|
||||
if request.TransitionFrames == 0 {
|
||||
request.TransitionFrames = 5
|
||||
}
|
||||
if request.TransitionFrames < 1 || request.TransitionFrames > 60 || request.Steps < 1 || request.Steps > 1000 {
|
||||
http.Error(w, "transition frames must be 1..60 and steps 1..1000", 400)
|
||||
return
|
||||
}
|
||||
if request.Model == "" {
|
||||
@ -229,7 +394,12 @@ func main() {
|
||||
http.Error(w, "selected motion model is not available: "+model.Reason, http.StatusConflict)
|
||||
return
|
||||
}
|
||||
a := &animation{ID: token(), Prompt: request.Prompt, Frames: request.Frames, DiffusionSteps: request.Steps, Seed: request.Seed, CreatedAt: time.Now().UTC().Format(time.RFC3339), Status: "queued", Kind: "generated", Model: request.Model}
|
||||
totalFrames := 0
|
||||
for _, segment := range request.Segments {
|
||||
totalFrames += segment.Frames
|
||||
}
|
||||
totalFrames -= request.TransitionFrames * (len(request.Segments) - 1)
|
||||
a := &animation{ID: token(), Prompt: request.Segments[0].Prompt, Frames: totalFrames, DiffusionSteps: request.Steps, Seed: request.Seed, CreatedAt: time.Now().UTC().Format(time.RFC3339), Status: "queued", Kind: "generated", Model: request.Model, Segments: request.Segments, TransitionFrames: request.TransitionFrames}
|
||||
g.mu.Lock()
|
||||
g.items[a.ID] = a
|
||||
err := g.save(a)
|
||||
|
||||
49
demo/main_test.go
Normal file
49
demo/main_test.go
Normal file
@ -0,0 +1,49 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestStitchSegmentsBlendsOverlap(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
first, second := filepath.Join(dir, "first"), filepath.Join(dir, "second")
|
||||
if err := os.MkdirAll(first, 0755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := os.MkdirAll(second, 0755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
rot := make([]float32, 2*22*4)
|
||||
for index := range rot {
|
||||
if index%4 == 3 {
|
||||
rot[index] = 1
|
||||
}
|
||||
}
|
||||
if err := writeF32(filepath.Join(first, "root_positions.f32"), []float32{0, 0, 0, 1, 0, 0}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := writeF32(filepath.Join(first, "local_rotations_xyzw.f32"), rot); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := writeF32(filepath.Join(second, "root_positions.f32"), []float32{3, 0, 0, 5, 0, 0}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := writeF32(filepath.Join(second, "local_rotations_xyzw.f32"), rot); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := stitchSegments(dir, []string{first, second}, 1); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
root, err := readF32(filepath.Join(dir, "root_positions.f32"))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(root) != 9 {
|
||||
t.Fatalf("frames = %d, want 3", len(root)/3)
|
||||
}
|
||||
if root[3] != 2 {
|
||||
t.Fatalf("blended boundary = %v, want 2", root[3])
|
||||
}
|
||||
}
|
||||
@ -1,36 +1,64 @@
|
||||
// Add the model picker without duplicating the compact canvas viewer script.
|
||||
// Extend the compact canvas viewer with model selection and a prompt sequence.
|
||||
window.addEventListener('load', async () => {
|
||||
const prompt = document.querySelector('#prompt');
|
||||
const form = prompt?.closest('.promptbox');
|
||||
if (!prompt || !form) return;
|
||||
const generate = document.querySelector('#generate');
|
||||
if (!prompt || !form || !generate) return;
|
||||
const models = await fetch('/api/models').then(r => r.json());
|
||||
const label = document.createElement('label');
|
||||
label.htmlFor = 'motionModel'; label.textContent = 'Motion model';
|
||||
|
||||
const modelLabel = document.createElement('label');
|
||||
modelLabel.htmlFor = 'motionModel'; modelLabel.textContent = 'Motion model';
|
||||
const select = document.createElement('select');
|
||||
select.id = 'motionModel';
|
||||
select.id = 'motionModel'; select.style.cssText = 'width:100%;padding:10px;border-radius:10px;background:#0b1015;color:#f3f6f4;border:1px solid #25313b';
|
||||
for (const model of models) {
|
||||
const option = document.createElement('option');
|
||||
option.value = model.id;
|
||||
const option = document.createElement('option'); option.value = model.id;
|
||||
option.disabled = !model.available;
|
||||
option.textContent = `${model.label}${model.available ? '' : ' — coming soon'}`;
|
||||
select.append(option);
|
||||
}
|
||||
const hint = document.createElement('div'); hint.className = 'hint';
|
||||
const update = () => {
|
||||
const modelHint = document.createElement('div'); modelHint.className = 'hint';
|
||||
const updateModel = () => {
|
||||
const model = models.find(item => item.id === select.value);
|
||||
hint.textContent = model.available
|
||||
? `${model.skeleton} · ${model.upstream}`
|
||||
: `${model.skeleton} · ${model.reason}`;
|
||||
modelHint.textContent = model.available ? `${model.skeleton} · ${model.upstream}` : `${model.skeleton} · ${model.reason}`;
|
||||
};
|
||||
select.onchange = update;
|
||||
form.insertBefore(label, prompt); form.insertBefore(select, prompt); form.insertBefore(hint, prompt);
|
||||
update();
|
||||
select.onchange = updateModel;
|
||||
form.insertBefore(modelLabel, prompt); form.insertBefore(select, prompt); form.insertBefore(modelHint, prompt); updateModel();
|
||||
|
||||
const sequence = document.createElement('div');
|
||||
sequence.style.cssText = 'display:grid;gap:10px;width:100%';
|
||||
prompt.before(sequence); sequence.append(prompt);
|
||||
prompt.classList.add('sequence-prompt');
|
||||
const segmentControls = new Map();
|
||||
const count = document.createElement('div'); count.className = 'hint';
|
||||
const updateCount = () => {
|
||||
const prompts = sequence.querySelectorAll('.sequence-prompt');
|
||||
count.textContent = `${prompts.length} segment${prompts.length === 1 ? '' : 's'} · 5-frame overlap`;
|
||||
};
|
||||
const addSegment = (text = '', frames = 150) => {
|
||||
const row = document.createElement('div'); row.style.cssText = 'display:grid;grid-template-columns:1fr 74px auto;gap:7px;align-items:start';
|
||||
const textArea = document.createElement('textarea'); textArea.className = 'sequence-prompt'; textArea.value = text;
|
||||
textArea.placeholder = 'Describe the next motion'; textArea.style.minHeight = '64px';
|
||||
const duration = document.createElement('input'); duration.type = 'number'; duration.min = '60'; duration.max = '300'; duration.step = '30'; duration.value = String(frames); duration.title = 'Frames (60–300)';
|
||||
const remove = document.createElement('button'); remove.type = 'button'; remove.textContent = '×'; remove.title = 'Remove segment'; remove.style.cssText = 'padding:8px 12px;background:#24313a;color:#dce9e8';
|
||||
remove.onclick = () => { row.remove(); updateCount(); };
|
||||
row.append(textArea, duration, remove); sequence.append(row); segmentControls.set(row, duration);
|
||||
updateCount();
|
||||
};
|
||||
const add = document.createElement('button'); add.type = 'button'; add.textContent = '+ Add prompt segment';
|
||||
add.style.cssText = 'justify-self:start;padding:8px 12px;background:#24313a;color:#dce9e8';
|
||||
add.onclick = () => addSegment(); form.insertBefore(add, generate); form.insertBefore(count, generate); updateCount();
|
||||
|
||||
const nativeFetch = window.fetch.bind(window);
|
||||
window.fetch = (input, init) => {
|
||||
if (typeof input === 'string' && input.endsWith('/api/generate') && init?.body) {
|
||||
const body = JSON.parse(init.body);
|
||||
body.model = select.value;
|
||||
body.transition_frames = 5;
|
||||
body.segments = [...sequence.querySelectorAll('.sequence-prompt')].map(area => {
|
||||
const row = area.closest('div');
|
||||
const duration = segmentControls.get(row);
|
||||
return {prompt: area.value, frames: Number(duration?.value || 150)};
|
||||
});
|
||||
return nativeFetch(input, {...init, body: JSON.stringify(body)});
|
||||
}
|
||||
return nativeFetch(input, init);
|
||||
|
||||
Loading…
Reference in New Issue
Block a user