【问题标题】:How to use C# async/await as a stand-alone CPS transform如何使用 C# async/await 作为独立的 CPS 转换
【发布时间】:2019-10-03 21:35:48
【问题描述】:

注意 1:这里 CPS 代表"continuation passing style"

我对了解如何连接到 C# 异步机制非常感兴趣。 基本上,据我了解 C# async/await 功能,编译器正在执行 CPS 转换,然后将转换后的代码传递给一个上下文对象,该对象管理各种线程上的任务调度。

您认为可以利用该编译器功能来创建 强大的组合器,同时抛开默认的线程方面?

一个例子是可以去递归和记忆像

这样的方法
async MyTask<BigInteger> Fib(int n)     // hypothetical example
{
    if (n <= 1) return n;
    return await Fib(n-1) + await Fib(n-2);
}

我设法做到了:

void Fib(int n, Action<BigInteger> Ret, Action<int, Action<BigInteger>> Rec)
{
    if (n <= 1) Ret(n);
    else Rec(n-1, x => Rec(n-2, y => Ret(x + y)));
}

(不使用异步,非常笨拙...)

或使用monad (While&lt;X&gt; = Either&lt;X, While&lt;X&gt;&gt;)

While<X> Fib(int n) => n <= 1 ?
    While.Return((BigInteger) n) :
    from x in Fib(n-1)
    from y in Fib(n-2)
    select x + y;

稍微好一点,但不像异步语法那么可爱:)


我在the blog of E. Lippert 上问过这个问题,他很友好地告诉我这确实是可能的。


在实现 ZBDD 库时需要我:(一种特殊的 DAG)

  • 大量复杂的相互递归操作

  • 实际示例中堆栈不断溢出

  • 只有在完全记忆的情况下才实用

手动 CPS 和去递归非常繁琐且容易出错。


对我所追求的东西(堆栈安全性)的严格测试是这样的:

async MyTask<BigInteger> Fib(int n, BigInteger a, BigInteger b)
{
    if (n == 0) return b;
    if (n == 1) return a;
    return await Fib(n - 1, a + b, a);
}

默认行为会在Fib(10000, 1, 0) 上产生堆栈溢出。或者更好的是,使用开头的代码和 memoization 来计算 Fib(10000)

【问题讨论】:

  • IEnumerable&lt;T&gt;/IEnumerator&lt;T&gt;yield 不能满足您的需求吗?它实际上是 async/await 的解耦机制。
  • 有可能:IEnumerator&lt;T&gt; 在概念上类似于Maybe[(T, IEnumerator&lt;T&gt;)],尽管是有状态的。我还提到了一个 monad 构造 While&lt;X&gt; = Either&lt;X, While&lt;X&gt;&gt; 可以解决问题,但实际上我的问题是关于劫持编译器在 async/await 语句上执行的 CPS 转换。
  • 已经可以通过使用自定义等待程序扩展此机制而无需进入编译器级别。至于记忆,我有严重的怀疑,因为async 语义并不暗示它。
  • @DmytroMukalov 读到 link 似乎很有希望,尤其是那个:Await task1.OnCoroutine(crm1)
  • 我知道,但请注意 F# 的计算表达式比 C# 对查询理解所做的更灵活,因为这些模式是固定的,而 F# 允许添加你自己的 -- 加上 F# 的类型和泛型使“做事”通常更容易发挥作用。我不争辩你可以用 C# 做这种事情,而且这个问题本身是有效的,但退后一步考虑你是否使用了正确的工具永远不会除了纯粹的智力练习之外,这也是一件坏事。我特别提到了 F#,因为 C# 和 F# 可以很容易地混合使用,具有 .NET 的共同点。

标签: c# async-await stack-overflow combinators continuation-passing


【解决方案1】:

这是我的解决方案版本。它是堆栈安全的,不使用线程池,但有特定的限制。特别是它需要尾递归风格的方法,所以像Fib(n-1) + Fib(n-2) 这样的结构将不起作用。另一方面,实际上以迭代方式执行的尾递归性质不需要记忆,因为每次迭代都被调用一次。它没有边缘情况保护,但它更像是一个原型而不是最终解决方案:

public class RecursiveTask<T>
{
    private T _result;

    private Func<RecursiveTask<T>> _function;

    public T Result
    {
        get
        {
            var current = this;
            var last = current;

            do
            {
                last = current;
                current = current._function?.Invoke();
            } while (current != null);

            return last._result;
        }
    }

    private RecursiveTask(Func<RecursiveTask<T>> function)
    {
        _function = function;
    }

    private RecursiveTask(T result)
    {
        _result = result;
    }

    public static implicit operator RecursiveTask<T>(T result)
    {
        return new RecursiveTask<T>(result);
    }

    public static RecursiveTask<T> FromFunc(Func<RecursiveTask<T>> func) => new RecursiveTask<T>(func);
}

及用法:

class Program
{
    static RecursiveTask<int> Fib(int n, int a, int b)
    {
        if (n == 0) return a;
        if (n == 1) return b;

        return RecursiveTask<int>.FromFunc(() => Fib(n - 1, b, a + b));
    }

    static RecursiveTask<int> Factorial(int n, int a)
    {
        if (n == 0) return a;

        return RecursiveTask<int>.FromFunc(() => Factorial(n - 1, n * a));
    }


    static void Main(string[] args)
    {
        Console.WriteLine(Factorial(5, 1).Result);
        Console.WriteLine(Fib(100000, 0, 1).Result);
    }
}

请注意,重要的是返回一个包装循环调用的函数,而不是调用本身,以避免真正的递归。

更新 下面是另一个实现,它仍然不使用 CPS 变换,但允许使用接近代数递归的语义,即它支持函数内部的多个类似递归的调用,并且不需要函数是尾递归的。

public class RecursiveTask<T1, T2>
{
    private readonly Func<RecursiveTask<T1, T2>, T1, T2> _func;
    private readonly Dictionary<T1, RecursiveTask<T1, T2>> _allTasks;
    private readonly List<RecursiveTask<T1, T2>> _subTasks;
    private readonly RecursiveTask<T1, T2> _rootTask;
    private T1 _arg;
    private T2 _result;
    private int _runsCount;
    private bool _isCompleted;
    private bool _isEvaluating;

    private RecursiveTask(Func<RecursiveTask<T1, T2>, T1, T2> func)
    {
        _func = func;
        _allTasks = new Dictionary<T1, RecursiveTask<T1, T2>>();
        _subTasks = new List<RecursiveTask<T1, T2>>();
        _rootTask = this;
    }

    private RecursiveTask(Func<RecursiveTask<T1, T2>, T1, T2> func, T1 arg, RecursiveTask<T1, T2> rootTask) : this(func)
    {
        _arg = arg;
        _rootTask = rootTask;
    }

    public T2 Run(T1 arg)
    {
        if (!_isEvaluating)
            BuildTasks(arg);

        if (_isEvaluating)
            return EvaluateTasks(arg);

        return default;
    }

    public static RecursiveTask<T1, T2> Create(Func<RecursiveTask<T1, T2>, T1, T2> func)
    {
        return new RecursiveTask<T1, T2>(func);
    }

    private void AddSubTask(T1 arg)
    {
        if (!_allTasks.TryGetValue(arg, out RecursiveTask<T1, T2> subTask))
        {
            subTask = new RecursiveTask<T1, T2>(_func, arg, this);
            _allTasks.Add(arg, subTask);
            _subTasks.Add(subTask);
        }
    }

    private T2 Run()
    {
        if (!_isCompleted)
        {
            var runsCount = _rootTask._runsCount;
            _result = _func(_rootTask, _arg);
            _isCompleted = runsCount == _rootTask._runsCount;
        }
        return _result;
    }

    private void BuildTasks(T1 arg)
    {
        if (_runsCount++ == 0)
            _arg = arg;

        if (EqualityComparer<T1>.Default.Equals(_arg, arg))
        {
            Run();

            var processed = 0;
            var addedTasksCount = _subTasks.Count;
            while (processed < addedTasksCount)
            {
                for (var i = processed; i < addedTasksCount; i++, processed++)
                    _subTasks[i].Run();
                addedTasksCount = _subTasks.Count;
            }
            _isEvaluating = true;
        }
        else
            AddSubTask(arg);
    }

    private T2 EvaluateTasks(T1 arg)
    {
        if (EqualityComparer<T1>.Default.Equals(_arg, arg))
        {
            foreach (var task in Enumerable.Reverse(_subTasks))
                task.Run();

            return Run();
        }
        else
        {
            if (_allTasks.TryGetValue(arg, out RecursiveTask<T1, T2> task))
                return task._isCompleted ? task._result : task.Run();
            else
                return default;
        }
    }
}

用法:

class Program
{
    static int Fib(int num)
    {
        return RecursiveTask<int, int>.Create((t, n) =>
        {
            if (n == 0) return 0;
            if (n == 1) return 1;

            return t.Run(n - 1) + t.Run(n - 2);
        }).Run(num);
    }

    static void Main(string[] args)
    {
        Console.WriteLine(Fib(7));
        Console.WriteLine(Fib(100000));
    }
}

作为好处,它是堆栈安全的,不使用线程池,没有async await 基础设施的负担,使用记忆并允许使用或多或少的可读语义。当前的实现意味着仅使用带有单个参数的函数。为了使其适用于更广泛的功能,应该为不同的泛型参数集提供类似的实现:

RecursiveTask<T1, T2, T3>
RecursiveTask<T1, T2, T3, T4>
...

【讨论】:

  • 您最后的话也是使Task.Run(() =&gt; _) 在GSerg 的解决方案中无堆栈的原因。
  • 这是你在这里的蹦床模式的一个很好的实现(Eric 在评论中也提到了它)。使用 CPS 转换,所有调用都成为有效的尾调用,因此您的解决方案将适用。我的问题的想法是编译器有效地进行了 CPS 转换以实现 async/await 语法,能够单独使用它并将其应用于您的蹦床不是很好吗? +1 用于研究工作,顺便说一下干净的可重用模式实现。
  • @SamuelVidal,我认为将蹦床和async await 机制结合起来是有问题的,因为蹦床意味着更窄的场景集。但是,我提供了另一个实现,它允许使用与您的第一个解决方案中使用的语义接近的语义。我认为使用async await 可以实现类似的效果,但我不确定这是否合理,因为可以在没有async 机制负担和复杂性的情况下完成类似的事情。
【解决方案2】:

对我所追求的东西(堆栈安全性)的严格测试是这样的:

async MyTask<BigInteger> Fib(int n, BigInteger a, BigInteger b)
{
    if (n == 0) return b;
    if (n == 1) return a;
    return await Fib(n - 1, a + b, a);
}

那岂不是很简单

public static Task<BigInteger> Fib(int n, BigInteger a, BigInteger b)
{
    if (n == 0) return Task.FromResult(b);
    if (n == 1) return Task.FromResult(a);

    return Task.Run(() => Fib(n - 1, a + b, a));
}

?


或者,不使用线程池,

public static async Task<BigInteger> Fib(int n, BigInteger a, BigInteger b)
{
    if (n == 0) return b;
    if (n == 1) return a;

    return await Task.FromResult(a + b).ContinueWith(t => Fib(n - 1, t.Result, a), TaskScheduler.FromCurrentSynchronizationContext()).Unwrap();
}

,除非我严重误解了什么。

【讨论】:

  • 我相信 Samuel 正在寻找一种对线程池工作线程不起作用的解决方案。我认为他正在寻找更像蹦床的东西。
  • 也可以使用Task.FromResult 代替Task.Run
  • @DmytroMukalov:那怎么会是无堆栈的呢?
  • @EricLippert 第二个版本适用于我的上下文线程。出于某种原因,我并不完全确定它是否正确,即使它看起来是正确的。
  • 这确实是堆栈安全的,并且具有快速和干净的优点。但我想看看是否可以在不更改初始代码的情况下完成。
【解决方案3】:

如果不查看您的 MyTask&lt;T&gt; 并查看该异常的堆栈跟踪,就不可能知道发生了什么。

看起来你要找的是Generalized async return types

您可以浏览the source 以了解ValueTaskValueTask&lt;T&gt; 的处理方式。

【讨论】:

  • 很容易判断发生了什么,该函数是递归的,它会调用自身 10000 次导致堆栈溢出。
  • 不错的建议,我已经查看了源代码。 (很重的东西^^=)
  • 这是编译器需要的。
【解决方案4】:

以下是更接近我所追求但尚未完全令人满意的解决方案。 它基于 GSerg 提出的堆栈安全解决方案的见解,并添加了记忆。

Pro 算法的核心(FibAux 方法使用干净的 async/await 语法)。

缺点它仍在使用线程池执行。

    // Core algorithm using the cute async/await syntax
    // (n.b. this would be exponential without memoization.)
    private static async Task<BigInteger> FibAux(int n)
    {
        if (n <= 1) return n;
        return await Rec(n - 1) + await Rec(n - 2);
    }

    public static Func<int, Task<BigInteger>> Rec { get; }
        = Utils.StackSafeMemoize<int, BigInteger>(FibAux);

    public static BigInteger Fib(int n)
        => FibAux(n).Result;

    [Test]
    public void Test()
    {
        Console.WriteLine(Fib(100000));
    }

    public static class Utils
    {
        // the combinator (still using the thread pool for execution)
        public static Func<X, Task<Y>> StackSafeMemoize<X, Y>(Func<X, Task<Y>> func)
        {
            var memo = new Dictionary<X, Y>();
            return x =>
            {
                Y result;
                if (!memo.TryGetValue(x, out result))
                {
                    return Task.Run(() => func(x).ContinueWith(task =>
                    {
                        var y = task.Result;
                        memo[x] = y;
                        return y;
                    }));
                }

                return Task.FromResult(result);
            };
        }
    } 

为了比较,这是不使用 async/await 的 cps 版本。


    public static BigInteger Fib(int n)
    {
        var fib = Memo<int, BigInteger>((m, rec, cont) =>
        {
            if (m <= 1) cont(m);
            else rec(m - 1, x => rec(m - 2, y => cont(x + y)));
        });

        return fib(n);
    }

    [Test]
    public void Test()
    {
        Console.WriteLine(Fib(100000));
    }

    // ---------

    public static Func<X, Y> Memo<X, Y>(Action<X, Action<X, Action<Y>>, Action<Y>> func)
    {
        var memo = new Dictionary<X, Y>(); // can be a Lru cache
        var stack = new Stack<Action>();

        Action<X, Action<Y>> rec = null;
        rec = (x, cont) =>
        {
            stack.Push(() =>
            {
                Y res;
                if (memo.TryGetValue(x, out res))
                {
                    cont(res);
                }
                else
                {
                    func(x, rec, y =>
                    {
                        memo[x] = y;
                        cont(y);
                    });
                }
            });
        };

        return x =>
        {
            var res = default(Y);
            rec(x, y => res = y);
            while (stack.Count > 0)
            {
                var next = stack.Pop();
                next();
            }

            return res;
        };
    }

【讨论】:

    猜你喜欢
    • 2021-10-19
    • 2021-04-10
    • 2019-12-14
    • 2014-02-16
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    相关资源
    最近更新 更多