|
|
|
|
|
|
|
|
|
|
|
|
|
|
| section "data" {
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| scheduler_janet_array:
|
| bits32[32] {0,0,0,0, 0,0,0,0, 0,0,0,0,0,0,0,0, 0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0};
|
| }
|
|
|
|
|
| foreign import ccall "sov_kv_allocate_blocks"
|
| sov_kv_allocate_blocks :: Word64 -> Word32 -> IO ();
|
|
|
| foreign import ccall "sov_cuda_flash_attention"
|
| sov_cuda_flash_attention :: Word32 -> Word32 -> Word64 -> Word64 -> Word64 -> Word64 -> Word64 -> Word64 -> Word32 -> Word32 -> IO ();
|
|
|
| foreign import ccall "sov_worm_checkpoint"
|
| sov_worm_checkpoint :: Word64 -> IO ();
|
|
|
| foreign import ccall "sov_worm_restore"
|
| sov_worm_restore :: Word64 -> IO ();
|
|
|
| foreign import ccall "sov_kv_append_tokens"
|
| sov_kv_append_tokens :: Word64 -> Word64 -> Word64 -> Word64 -> Word32 -> Word32 -> IO ();
|
|
|
| foreign import ccall "sov_speculative_draft"
|
| sov_speculative_draft :: Word64 -> Word32 -> Word64 -> IO ();
|
|
|
| foreign import ccall "sov_speculative_verify"
|
| sov_speculative_verify :: Word64 -> Word64 -> Word32 -> IO Word32;
|
|
|
| foreign import ccall "sov_power_event"
|
| sov_power_event :: Word32 -> IO ();
|
|
|
|
|
|
|
|
|
| scheduler_step(W_ state, W_ batch_ptr, W_ kv_ptr)
|
| {
|
| W_ pending, new_state, tokens_gen, power;
|
|
|
| power = W_[scheduler_janet_array + (5 * SIZEOF_W)];
|
| if (power == 1) {
|
| jump scheduler_do_checkpoint(state, batch_ptr, kv_ptr);
|
| }
|
| if (power == 2) {
|
| jump scheduler_do_resume(state, batch_ptr, kv_ptr);
|
| }
|
|
|
| switch [0..5] state {
|
| case 0: {
|
| pending = W_[scheduler_janet_array + (0 * SIZEOF_W)];
|
| if (pending > 0) {
|
| new_state = 1;
|
| } else {
|
| new_state = 0;
|
| }
|
| return (new_state);
|
| }
|
| case 1: {
|
|
|
| foreign "C" sov_kv_allocate_blocks(kv_ptr, W_[scheduler_janet_array + (1 * SIZEOF_W)]);
|
|
|
| W_[scheduler_janet_array + (2 * SIZEOF_W)] = 0;
|
| new_state = 2;
|
| return (new_state);
|
| }
|
| case 2: {
|
| W_ heads, seqs, head_dim, block_size;
|
| seqs = W_[scheduler_janet_array + (1 * SIZEOF_W)];
|
| heads = 32;
|
| head_dim = 128;
|
| block_size = 16;
|
|
|
|
|
| foreign "C" sov_cuda_flash_attention(
|
| seqs, heads,
|
| batch_ptr,
|
| kv_ptr,
|
| kv_ptr + 4096,
|
| batch_ptr + 8192,
|
| kv_ptr + 16384,
|
| kv_ptr + 32768,
|
| head_dim, block_size);
|
|
|
|
|
| foreign "C" sov_speculative_draft(batch_ptr, W_[scheduler_janet_array + (6 * SIZEOF_W)], kv_ptr);
|
|
|
|
|
| tokens_gen = W_[scheduler_janet_array + (2 * SIZEOF_W)] + 1;
|
| W_[scheduler_janet_array + (2 * SIZEOF_W)] = tokens_gen;
|
|
|
|
|
| if ((tokens_gen & 63) == 0) {
|
| foreign "C" sov_worm_checkpoint(kv_ptr);
|
|
|
| W_[scheduler_janet_array + (7 * SIZEOF_W)] = W_[scheduler_janet_array + (7 * SIZEOF_W)] + 1;
|
| }
|
|
|
|
|
| if (tokens_gen >= 2048) {
|
|
|
| W_[scheduler_janet_array + (0 * SIZEOF_W)] = W_[scheduler_janet_array + (0 * SIZEOF_W)] - 1;
|
| new_state = 0;
|
| } else {
|
| new_state = 2;
|
| }
|
| return (new_state);
|
| }
|
| case 3: {
|
|
|
| foreign "C" sov_worm_checkpoint(kv_ptr);
|
| new_state = 1;
|
| return (new_state);
|
| }
|
| case 4: {
|
| jump scheduler_do_checkpoint(state, batch_ptr, kv_ptr);
|
| }
|
| case 5: {
|
| jump scheduler_do_resume(state, batch_ptr, kv_ptr);
|
| }
|
| default: {
|
| new_state = 0;
|
| return (new_state);
|
| }
|
| }
|
| }
|
|
|
| scheduler_do_checkpoint(W_ state, W_ batch_ptr, W_ kv_ptr)
|
| {
|
| foreign "C" sov_worm_checkpoint(kv_ptr);
|
|
|
| W_[scheduler_janet_array + (5 * SIZEOF_W)] = 0;
|
| return (5);
|
| }
|
|
|
| scheduler_do_resume(W_ state, W_ batch_ptr, W_ kv_ptr)
|
| {
|
| foreign "C" sov_worm_restore(kv_ptr);
|
| W_[scheduler_janet_array + (5 * SIZEOF_W)] = 0;
|
| return (1);
|
| }
|
|
|
|
|
|
|
|
|
| janet_get(W_ slot)
|
| {
|
| W_ val;
|
| val = W_[scheduler_janet_array + (slot * SIZEOF_W)];
|
| return (val);
|
| }
|
|
|
| janet_set(W_ slot, W_ val)
|
| {
|
| W_[scheduler_janet_array + (slot * SIZEOF_W)] = val;
|
| return ();
|
| }
|
|
|