Editor-side generator that wraps the UniMate model behind IMotionGenerator. Prepares rigs, runs text-to-motion, in-between, edit, expansion and variation modes by sampling the UniMate model in overlapping windows, assembling and postprocessing produced frames.
using System;
using System.Collections.Generic;
using System.Linq;
using System.Numerics;
using System.Threading;
using System.Threading.Tasks;
using TextToAnimation.Animation;
using TextToAnimation.Generation;
using TextToAnimation.Maths;
using TextToAnimation.Processing;
using TextToAnimation.Workspace;
namespace TextToAnimation.Editor.Inference.UniMate;
/// <summary>
/// UniMate behind the editor's <see cref="IMotionGenerator"/> interface. All modes run on 2-second windows
/// (UniMate's native 60 frames at 30 fps); longer clips are covered by overlapping windows, each one encoded
/// from the current result so the overlap is held fixed and the root path continues smoothly.
/// </summary>
public sealed class UniMateGenerator : IMotionGenerator
{
const int Window = UniMateModel.Frames;
const int Overlap = 10;
readonly UniMateModel _model;
readonly Dictionary<string, (UniMateRig Rig, PreparedSkeleton Prep)> _rigs = new();
readonly object _lock = new();
public UniMateGenerator( UniMateModel model ) => _model = model;
public string Name => "UniMate";
public GeneratorCapabilities Capabilities { get; } = new()
{
Modes = new[] { GenerationMode.TextToMotion, GenerationMode.InBetween, GenerationMode.TextEdit, GenerationMode.Expansion, GenerationMode.Variation },
NativeFps = UniMateModel.Fps,
MaxSegmentSeconds = Window / UniMateModel.Fps,
DefaultSeconds = 2f,
DefaultGuidance = 3f,
};
public IReadOnlyList<string> Validate( MotionRig rig ) => UniMateRig.Validate( rig );
public RigFamily DetectFamily( MotionRig rig ) => UniMateRig.DetectFamily( rig );
(UniMateRig Rig, PreparedSkeleton Prep) Prepare( MotionRig rig, RigFamily family, CancellationToken token )
{
if ( family == RigFamily.Auto ) family = UniMateRig.DetectFamily( rig );
var key = AnimationWorkspace.Fingerprint( rig.Skeleton ) + rig.Skeleton.RestWorld.Sum( x => x.Pos.X + x.Pos.Y * 3 + x.Pos.Z * 7 ).ToString( "R" ) + "|" + family;
lock ( _lock )
{
if ( _rigs.TryGetValue( key, out var cached ) ) return cached;
var uniRig = UniMateRig.Build( rig, family );
var prep = _model.Prepare( uniRig.Skeleton, UniMateStats.For( uniRig.Family ), token );
return _rigs[key] = (uniRig, prep);
}
}
public Task<IReadOnlyList<GeneratedMotion>> GenerateAsync( MotionRig rig, GenerationRequest request,
Action<GenerationProgress> progress, CancellationToken token )
=> Task.Run( () => Generate( rig, request, progress, token ), token );
IReadOnlyList<GeneratedMotion> Generate( MotionRig rig, GenerationRequest request, Action<GenerationProgress> progress, CancellationToken token )
{
progress?.Invoke( new GenerationProgress( "Preparing the skeleton", 0f ) );
var (uniRig, prep) = Prepare( rig, request.RigFamily, token );
var takes = Math.Max( 1, request.Count );
var steps = request.Steps > 0 ? request.Steps : 24;
var guidance = request.Guidance > 0 ? request.Guidance : 3f;
var results = new List<GeneratedMotion>();
// source motion at 30 fps (for in-betweening, editing and variations)
List<XForm[]> source = null;
var pins = new List<int>();
if ( request.SourceFrames is { Count: > 1 } src )
{
var scale = UniMateModel.Fps / Math.Max( 1f, request.SourceFps );
var count = Math.Max( 2, (int)MathF.Round( (src.Count - 1) * scale ) + 1 );
source = new List<XForm[]>( count );
for ( var f = 0; f < count; f++ ) source.Add( ClipOps.Sample( src, f / scale ) );
pins = request.KeepFrames.Select( f => (int)MathF.Round( f * scale ) ).Where( f => f >= 0 && f < count ).Distinct().OrderBy( f => f ).ToList();
}
for ( var take = 0; take < takes; take++ )
{
token.ThrowIfCancellationRequested();
var seed = request.Seed + take * 7919;
var notes = new List<string>();
List<XForm[]> frames;
var seams = new List<int>(); // model frames where a chained window begins
void Report( string stage, float fraction ) => progress?.Invoke( new GenerationProgress(
takes > 1 ? $"{stage} (take {take + 1} of {takes})" : stage, (take + Math.Clamp( fraction, 0, 1 )) / takes ) );
switch ( request.Mode )
{
case GenerationMode.TextToMotion:
case GenerationMode.Expansion:
{
// one prompt, one caption, as upstream (sample.py encodes the prompt whole; segments only come from a
// list of prompts, motion expansion): its captions say "runs, flips, and then lands" in one sentence
var segments = request.Prompts.Where( p => !string.IsNullOrWhiteSpace( p ) ).ToList();
if ( segments.Count == 0 ) throw new InvalidOperationException( "Describe the motion first." );
// a single prompt longer than one window repeats as continuing segments
var wanted = request.DurationSeconds > 0 ? request.DurationSeconds : Capabilities.DefaultSeconds;
while ( segments.Count * (Window - Overlap) + Overlap < wanted * UniMateModel.Fps && segments.Count < 20 )
segments.Add( segments[^1] );
frames = TextChain( uniRig, prep, segments.Select( UniMatePrompt.Caption ).ToList(), seed, steps, guidance, Report, token );
for ( var w = 1; w < segments.Count; w++ ) seams.Add( Window + (w - 1) * (Window - Overlap) );
if ( request.Mode == GenerationMode.TextToMotion && wanted > 0 )
{
var keep = Math.Clamp( (int)MathF.Round( wanted * UniMateModel.Fps ) + 1, 2, frames.Count );
frames = frames.GetRange( 0, keep );
}
break;
}
case GenerationMode.InBetween:
{
if ( source is null ) throw new InvalidOperationException( "In-betweening needs an animation with pinned frames." );
if ( pins.Count < 2 ) throw new InvalidOperationException( "Pin at least two frames (the poses to keep) on the timeline." );
var caption = UniMatePrompt.Caption( request.Prompts.FirstOrDefault() ?? "" );
frames = Chain( uniRig, prep, null, seed, steps, guidance, source, pins, null, 0f, Report, token, caption, seams );
break;
}
case GenerationMode.TextEdit:
{
if ( source is null ) throw new InvalidOperationException( "Editing needs an existing animation." );
var caption = UniMatePrompt.Caption( request.Prompts.FirstOrDefault() ?? "" );
if ( caption.Length == 0 ) throw new InvalidOperationException( "Describe the new motion for the unlocked bones." );
var keepJoints = uniRig.JointsForBones( request.KeepBones );
if ( keepJoints.Count == 0 ) notes.Add( "No bones were locked, so the whole body was regenerated." );
frames = Chain( uniRig, prep, null, seed, steps, guidance, source, null, keepJoints, 0f, Report, token, caption, seams );
break;
}
case GenerationMode.Variation:
{
if ( source is null ) throw new InvalidOperationException( "Variations need an existing animation." );
var caption = UniMatePrompt.Caption( request.Prompts.FirstOrDefault() ?? "An object moves." );
// start part-way along the flow from a noised copy of the source (SDEdit): low strength stays close
var t0 = Math.Clamp( 1f - request.VariationStrength, 0.05f, 0.9f );
frames = Chain( uniRig, prep, null, seed, steps, guidance, source, null, null, t0, Report, token, caption, seams );
break;
}
default:
throw new NotSupportedException( $"{request.Mode} is not supported." );
}
// derived bones: twist helpers follow their limbs. Bones UniMate animates are excluded - they already
// carry generated motion, and re-deriving them (the retargeter's "inline" follow) would override it and
// stretch the bones below them
TwistBoneFollow.Apply( frames, rig.Rig, uniRig.AnimatedBones );
// back to the workspace frame rate
var output = Resample( frames, UniMateModel.Fps, request.OutputFps );
if ( request.CleanUp )
{
var outputSeams = seams.Select( f => (int)MathF.Round( f * request.OutputFps / UniMateModel.Fps ) ).ToList();
var cleaned = ClipCleanup.CleanGenerated( output, rig, request.OutputFps, outputSeams );
if ( cleaned.Length > 0 ) notes.Add( cleaned );
}
EnforceConstraints( output, request );
results.Add( new GeneratedMotion { Frames = output, Fps = request.OutputFps, Seed = seed, Notes = notes } );
}
progress?.Invoke( new GenerationProgress( "Done", 1f ) );
return results;
}
/// <summary>
/// Text to motion as upstream generates it (sample.py with motion expansion): the first window freely (dopri5),
/// each further window with its first <see cref="Overlap"/> frames held to the previous window's last ones
/// (replacement sampling, in the network's feature space), the windows' features joined at the seams and decoded
/// once - so the root path and the poses run on continuously across windows.
/// </summary>
List<XForm[]> TextChain( UniMateRig uniRig, PreparedSkeleton prep, List<string> captions, int seed, int steps, float guidance,
Action<string, float> report, CancellationToken token )
{
var J = uniRig.Count;
const int T = Window;
var stats = UniMateStats.For( uniRig.Family );
var windows = captions.Count;
var total = T + (windows - 1) * (T - Overlap);
var feat = new float[total, J, 12];
float[] previous = null;
for ( var w = 0; w < windows; w++ )
{
token.ThrowIfCancellationRequested();
var embedding = _model.Text.Encode( captions[w], token );
var noise = UniMateModel.Noise( J, seed + w * 104729 );
SampleSettings settings;
if ( previous is null ) settings = FreeSampler( steps, guidance );
else
{
var known = new float[J * 12 * T];
var keep = new bool[J * 12 * T];
for ( var jc = 0; jc < J * 12; jc++ )
for ( var f = 0; f < Overlap; f++ )
{
known[jc * T + f] = previous[jc * T + T - Overlap + f];
keep[jc * T + f] = true;
}
settings = ConstrainedSampler( steps, guidance, known, keep, 0f );
}
var windowIndex = w;
var x = _model.Sample( prep, embedding, noise, settings,
f => report( windows > 1 ? $"Generating part {windowIndex + 1} of {windows}" : "Generating", (windowIndex + f) / windows ), token );
previous = x;
var wf = UniMateFeatures.FromModel( x, J, T, stats );
var start = w == 0 ? 0 : Overlap;
var at = w == 0 ? 0 : T + (w - 1) * (T - Overlap);
for ( var f = start; f < T; f++ )
for ( var j = 0; j < J; j++ )
for ( var c = 0; c < 12; c++ )
feat[at + f - start, j, c] = wf[f, j, c];
}
var motion = UniMateFeatures.Decode( feat, uniRig.Skeleton.Parents );
return uniRig.ToFrames( UniMateFeatures.ToSource( motion, uniRig.Skeleton ), null );
}
/// <summary>
/// Free sampling as upstream samples (Sampler.sample_ode: dopri5 integrated to convergence, rtol 1e-3, atol 1e-6).
/// "Fast" loosens the tolerance (rtol 1e-2: within 0.004 of upstream's motion on a body of diameter 2, a few
/// dozen calls fewer). Fixed-step shortcuts are not offered: 16 Adams-Bashforth calls land 0.46 away from
/// upstream's motion, 50 Euler steps 0.13 (DatasetMotionTests, T2A_SAMPLER_STUDY).
/// </summary>
static SampleSettings FreeSampler( int steps, float guidance ) => steps <= 12
? new SampleSettings { Method = Integrator.Dopri5, Guidance = guidance, RelativeTolerance = 1e-2, AbsoluteTolerance = 1e-5 }
: new SampleSettings { Method = Integrator.Dopri5, Guidance = guidance };
/// <summary>
/// Sampling with held values (in-betweening, editing, window seams) as upstream does it (inbetween_sample_ode):
/// fixed-step Euler with replacement, 50 steps over the flow (fewer when a variation starts part-way along it).
/// </summary>
static SampleSettings ConstrainedSampler( int steps, float guidance, float[] known, bool[] keep, float startTime ) =>
new() { Method = Integrator.Euler, Steps = Math.Max( 8, (int)MathF.Round( 50 * (1 - startTime) ) ), Guidance = guidance, Known = known, Keep = keep, StartTime = startTime };
/// <summary>
/// Generates window after window. Each window is encoded from the current result (source frames, or the
/// motion generated so far) and holds: pinned frames, kept joints, and the overlap with the previous window.
/// </summary>
List<XForm[]> Chain( UniMateRig uniRig, PreparedSkeleton prep, List<string> captions, int seed, int steps, float guidance,
List<XForm[]> source, List<int> pins, HashSet<int> keepJoints, float startTime,
Action<string, float> report, CancellationToken token, string singleCaption = null, List<int> seams = null )
{
var J = uniRig.Count;
var length = source?.Count ?? (Window + (captions.Count - 1) * (Window - Overlap) + 1);
var windows = new List<int>();
for ( var s = 0; ; s += Window - Overlap )
{
windows.Add( s );
if ( s + Window >= length - 1 ) break;
}
var result = source is not null ? AnimClip.CopyFrames( source ) : new List<XForm[]>();
var rest = uniRig.Motion.Skeleton.Bones.Select( b => b.RestLocal ).ToArray();
for ( var w = 0; w < windows.Count; w++ )
{
token.ThrowIfCancellationRequested();
var start = windows[w];
var caption = singleCaption ?? captions[Math.Min( w, captions.Count - 1 )];
var embedding = _model.Text.Encode( caption, token );
var noise = UniMateModel.Noise( J, seed + w * 104729 );
var keep = new bool[J * 12 * Window];
var anyKeep = false;
void KeepFrame( int f ) { for ( var j = 0; j < J; j++ ) for ( var c = 0; c < 12; c++ ) keep[(j * 12 + c) * Window + f] = true; anyKeep = true; }
if ( pins is not null ) foreach ( var p in pins ) if ( p >= start && p < start + Window ) KeepFrame( p - start );
if ( keepJoints is not null )
foreach ( var j in keepJoints ) for ( var c = 0; c < 12; c++ ) for ( var f = 0; f < Window; f++ ) { keep[(j * 12 + c) * Window + f] = true; anyKeep = true; }
if ( w > 0 ) for ( var f = 0; f < Overlap; f++ ) KeepFrame( f );
if ( w > 0 ) seams?.Add( start + Overlap ); // the first frame this window adds
float[] known = null;
UniMateFeatures.Alignment alignment = null;
var needsKnown = anyKeep || startTime > 0;
if ( needsKnown )
{
// T+1 frames of the current result for this window (padded with the last available frame)
var span = new List<XForm[]>( Window + 1 );
for ( var f = 0; f <= Window; f++ )
{
var i = Math.Min( start + f, result.Count - 1 );
span.Add( result.Count > 0 ? result[Math.Max( 0, i )] : rest );
}
var (pos, rot) = uniRig.JointWorld( span );
var (feat, align) = UniMateFeatures.Encode( pos, rot, uniRig.Skeleton );
alignment = align;
known = UniMateFeatures.ToModel( feat, UniMateStats.For( uniRig.Family ), Window );
}
var settings = ConstrainedSampler( steps, guidance, known, anyKeep ? keep : null, startTime );
var windowIndex = w;
var x = _model.Sample( prep, embedding, noise, settings,
f => report( windows.Count > 1 ? $"Generating part {windowIndex + 1} of {windows.Count}" : "Generating", (windowIndex + f) / windows.Count ), token );
var motion = UniMateFeatures.Decode( UniMateFeatures.FromModel( x, J, Window, UniMateStats.For( uniRig.Family ) ), uniRig.Skeleton.Parents );
if ( alignment is not null ) UniMateFeatures.Unalign( motion, alignment );
var src = UniMateFeatures.ToSource( motion, uniRig.Skeleton );
// bones UniMate doesn't animate (fingers, helpers) keep the source pose
var baseFrames = source is not null ? Enumerable.Range( 0, Window ).Select( f => source[Math.Min( start + f, source.Count - 1 )] ).ToList() : null;
var generated = uniRig.ToFrames( src, baseFrames );
// write the window into the result
for ( var f = 0; f < Window; f++ )
{
var target = start + f;
if ( target < result.Count ) result[target] = generated[f];
else result.Add( generated[f] );
}
}
if ( source is not null && result.Count > source.Count ) result = result.GetRange( 0, source.Count );
return result;
}
/// <summary>Pinned poses (in-betweening) and locked bones (editing) come back exactly as they were.</summary>
static void EnforceConstraints( List<XForm[]> output, GenerationRequest request )
{
if ( request.SourceFrames is not { Count: > 1 } src ) return;
var keepPins = request.Mode == GenerationMode.InBetween && request.KeepFrames.Count > 0;
var keepBones = request.Mode == GenerationMode.TextEdit && request.KeepBones.Count > 0;
if ( !keepPins && !keepBones ) return;
// the source on the output's frame grid
var scale = request.OutputFps / Math.Max( 1f, request.SourceFps );
IReadOnlyList<XForm[]> source = src;
if ( MathF.Abs( scale - 1f ) > 1e-4f )
{
var grid = new List<XForm[]>( output.Count );
for ( var f = 0; f < output.Count; f++ ) grid.Add( ClipOps.Sample( src, f / scale ) );
source = grid;
}
if ( keepBones ) GenerationConstraints.RestoreBones( output, source, request.KeepBones );
if ( keepPins ) GenerationConstraints.RestorePins( output, source, request.KeepFrames.Select( f => (int)MathF.Round( f * scale ) ) );
}
static List<XForm[]> Resample( List<XForm[]> frames, float fromFps, float toFps )
{
if ( MathF.Abs( fromFps - toFps ) < 0.01f || frames.Count < 2 ) return frames;
var duration = (frames.Count - 1) / fromFps;
var count = Math.Max( 2, (int)MathF.Round( duration * toFps ) + 1 );
var result = new List<XForm[]>( count );
for ( var f = 0; f < count; f++ ) result.Add( ClipOps.Sample( frames, f * fromFps / toFps ) );
return result;
}
}